Skip to content

Instantly share code, notes, and snippets.

@alisterburt
Last active March 24, 2023 20:58
Show Gist options
  • Select an option

  • Save alisterburt/622d462a55960bf6e4f28543d4fbbf19 to your computer and use it in GitHub Desktop.

Select an option

Save alisterburt/622d462a55960bf6e4f28543d4fbbf19 to your computer and use it in GitHub Desktop.
libtilt motion estimation first pass
import torch
import numpy as np
import mrcfile
import einops
import torch.nn.functional as F
from torch_cubic_spline_grids import CubicBSplineGrid3d
from libtilt.utils.coordinates import array_to_grid_sample
from libtilt.shift.phase_shift_2d import phase_shift_images_2d, phase_shift_dfts_2d
from libtilt.mask.shapes_2d import circle
from libtilt.patch.patch_indices import get_patch_indices_2d
from libtilt.patch.patch_centers import get_patch_centers_2d, get_patch_centers_3d
IMAGE_FILE = 'data/TS_01_000_0_0_bin8.mrc'
PATCH_SIDELENGTH = 64 # 64 * 8 = 512
GT_DEFORMATION_FIELD_RESOLUTION = (3, 3, 3) # (t, h, w)
LEARNED_DEFORMATION_FIELD_RESOLUTION = (5, 5, 5) # (t, h, w)
LEARNING_RATE = 0.9
image = torch.tensor(mrcfile.read(IMAGE_FILE))
image = (image - torch.mean(image)) / torch.std(image)
### simulate continuous beam induced motion on the image
multi_frame_micrograph = einops.repeat(image, 'h w -> 10 h w').clone()
t, h, w = multi_frame_micrograph.shape[-3:]
gt_deformation_field_data = np.random.uniform(
low=-5, high=5, size=(2, *GT_DEFORMATION_FIELD_RESOLUTION)
)
gt_deformation_field = CubicBSplineGrid3d.from_grid_data(
torch.tensor(gt_deformation_field_data).float()
)
_t, _y, _x = torch.linspace(0, 1, t), torch.linspace(0, 1, h), torch.linspace(0, 1, w)
tt, yy, xx = torch.meshgrid(_t, _y, _x)
tyx = einops.rearrange([tt, yy, xx], 'tyx t h w -> t h w tyx')
shifts = gt_deformation_field(tyx)
array_coordinates = torch.tensor(np.indices(image.shape[-2:]))
array_coordinates = einops.rearrange(array_coordinates, 'yx ... -> ... yx')
deformed_coordinates = array_coordinates + shifts
grid_sample_coordinates = array_to_grid_sample(deformed_coordinates, array_shape=(h, w))
multi_frame_micrograph = F.grid_sample(
input=einops.rearrange(multi_frame_micrograph, 't h w -> t 1 h w'),
grid=grid_sample_coordinates,
mode='bicubic',
padding_mode='zeros',
align_corners=True,
)
multi_frame_micrograph = einops.rearrange(multi_frame_micrograph, 't 1 h w -> t h w')
### learn the motion
reference = image.clone()
# extract patches...
patch_idx_h, patch_idx_w = get_patch_indices_2d(
image_shape=image.shape,
patch_shape=(PATCH_SIDELENGTH, PATCH_SIDELENGTH),
patch_step=(PATCH_SIDELENGTH // 2, PATCH_SIDELENGTH // 2),
distribute_patches=True,
)
data_patches = multi_frame_micrograph[:, patch_idx_h, patch_idx_w].detach()
reference_patches = reference[patch_idx_h, patch_idx_w].detach()
patch_centers = get_patch_centers_3d(
image_shape=multi_frame_micrograph.shape,
patch_shape=(1, PATCH_SIDELENGTH, PATCH_SIDELENGTH),
patch_step=(1, PATCH_SIDELENGTH // 2, PATCH_SIDELENGTH // 2),
distribute_patches=True
) # (t, h, w, thw)
patch_centers = patch_centers / torch.tensor([t - 1, h - 1, w - 1])
# mask the reference and the data
mask = circle(
radius=PATCH_SIDELENGTH / 4, sidelength=PATCH_SIDELENGTH,
smoothing_radius=PATCH_SIDELENGTH / 8
)
data_patches *= mask
reference_patches *= mask
deformation_field = CubicBSplineGrid3d(
resolution=LEARNED_DEFORMATION_FIELD_RESOLUTION,
n_channels=2
)
motion_optimiser = torch.optim.Adam(
params=deformation_field.parameters(),
lr=LEARNING_RATE,
)
for i in range(100):
motion_optimiser.zero_grad()
shifts = deformation_field(patch_centers)
image_shape = (PATCH_SIDELENGTH, PATCH_SIDELENGTH)
shifted_patches = phase_shift_images_2d(
data_patches,
shifts=shifts,
)
loss = torch.sqrt(torch.mean((reference_patches - shifted_patches) ** 2))
last_loss = loss.item()
loss.backward()
motion_optimiser.step()
print(loss.item())
# quantify how well we're doing
gt = gt_deformation_field(patch_centers)
learned = deformation_field(patch_centers)
print(torch.mean(torch.abs(gt - learned)))
# invert the motion
shifts = -1 * deformation_field(tyx)
sample_coords = array_coordinates + shifts
reconstruction = F.grid_sample(
input=einops.rearrange(multi_frame_micrograph, 't h w -> t 1 h w'),
grid=array_to_grid_sample(sample_coords, array_shape=(h, w)),
mode='bicubic',
padding_mode='zeros',
align_corners=True,
)
reconstruction = einops.rearrange(reconstruction, 't 1 h w -> t h w')
# visualise
import napari
viewer = napari.Viewer()
viewer.add_image(reference.detach().numpy(), name='reference')
viewer.add_image(multi_frame_micrograph.detach().numpy(), name='wavy data')
viewer.add_image(reconstruction.detach().numpy(), name='motion corrected')
napari.run()
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment