import skimage

import matplotlib.pyplot as plt
import numpy as np
import math

### Load frames and upscale
N = 9
imgs = []
for i in range(N):
    img = skimage.io.imread(f'frames/frame{i+1:01}.png')
    imgs.append(img)
    imgs[i] = skimage.transform.rescale(imgs[i], (2, 2, 1), anti_aliasing=False, order=1)
    print(f'Loaded Frame {i+1}/{N}.', end='\r')



### Return the projective transformation that best aligns 2 images using RANSAC
def get_orb_alignment_transformation(target, source):
    orb_target = skimage.feature.ORB()
    orb_target.detect_and_extract(skimage.color.rgb2gray(target))
    orb_source = skimage.feature.ORB()
    orb_source.detect_and_extract(skimage.color.rgb2gray(source))
    
    feature_matches = skimage.feature.match_descriptors(orb_target.descriptors, orb_source.descriptors, metric='hamming', cross_check=True)
    
    matched_target_coords = np.array([orb_target.keypoints[i[0]] for i in feature_matches])
    matched_source_coords = np.array([orb_source.keypoints[i[1]] for i in feature_matches])
    
    transform = skimage.measure.ransac((matched_source_coords, matched_target_coords), 
                                        skimage.transform.ProjectiveTransform, 
                                        min_samples=4,
                                        residual_threshold=0.71, 
                                        max_trials=1000,
                                        stop_probability=0.9999999)[0].params
    
    transform[[0,1],:] = transform[[1,0],:]
    transform[:,[0,1]] = transform[:,[1,0]]
    
    return transform, orb_target, orb_source, feature_matches



### Compute alignment transformation for each frame
alignment_transformations = []

img0 = imgs[0]
for i, img in enumerate(imgs):
    transform, *_ = get_orb_alignment_transformation(img0, img)
    alignment_transformations.append(np.linalg.inv(transform))
    print(f"Alignment transformation found {i+1}/{N}.", end='\r')



### Display the aligned frames
fig, axs = plt.subplots(nrows=3, ncols=3)
fig.set_figheight(20)
fig.set_figwidth(25)
for i, (img, transform) in enumerate(zip(imgs, alignment_transformations)):
    aligned_image = skimage.transform.warp(img, transform)
    axs[i//3][i%3].imshow(skimage.util.compare_images(img0, aligned_image, method='diff'))
    axs[i//3][i%3].set_title(f'Frame {i}')
plt.show()



### Some frames are badly aligned, so skip them
skip = [2, 6]

### Align the images and average them.
margin = 50
height, width, k = imgs[0].shape
out_shape = height + 2 * margin, width + 2 * margin
glob_trfm = np.eye(3)
glob_trfm[:2, 2] = -margin, -margin

aligned_images = [skimage.transform.warp(img, trfm.dot(glob_trfm),
                                         output_shape=out_shape,
                                         mode="constant", cval=np.nan)
                  for i, (img, trfm) in enumerate(zip(imgs, alignment_transformations)) if i not in skip]

combined = np.mean(aligned_images, 0)

skimage.io.imshow(combined)
skimage.io.imsave('combined.png', combined)