"""
====================================
Readout oversampling and apodization
====================================

Two operations are applied to Cartesian k-space before the image is
reconstructed. Readout oversampling doubles the field of view along the
frequency-encoding direction, and is removed so that the image has the
prescribed matrix. Truncation of k-space at the edge of the acquired matrix
convolves the image with a sinc, whose side lobes appear as Gibbs ringing
parallel to every sharp edge; an apodization window reduces the ringing at the
cost of spatial resolution.

This example removes twofold readout oversampling from the k-space of a
Shepp-Logan phantom and compares the result with an acquisition without
oversampling, then reconstructs a 64 x 64 acquisition with no window, a Fermi
window and a Hann window, and measures the ringing and the resolution of each.
The k-space is evaluated analytically at the sample positions, so its
truncation and its oversampling are those of a continuous object.

**Learning objectives**

* Remove readout oversampling in the image domain with
  :func:`bartorch.remove_readout_oversampling`, and distinguish it from
  discarding the outer readout samples.
* Relate the Gibbs ringing of a truncated acquisition to the side lobes of its
  point spread function.
* Apodize k-space with :func:`bartorch.apodize`, and quantify the trade-off
  between ringing amplitude and the full width at half maximum.
* Choose between the radial and the separable extension of a window over
  k-space.
"""

# %%

# sphinx_gallery_start_ignore
import matplotlib.pyplot as plt
from matplotlib.patches import Rectangle

BOX = "#3dbde8"
COLOURS = {"no window": "#e8a33d", "Fermi": "#3dbde8", "Hann": "#c65fd6"}


def show(axis, values, title, vmin=0.0, vmax=1.0, cmap="gray"):
    handle = axis.imshow(values, cmap=cmap, vmin=vmin, vmax=vmax)
    axis.set_title(title)
    axis.set_axis_off()
    return handle


# sphinx_gallery_end_ignore
import math

import torch

import bartorch
import bartorch.tools as bt

# %%
#
# Readout oversampling
# --------------------
#
# The receiver digitises the echo at twice the sampling rate the prescribed
# matrix requires, so that its anti-aliasing filter can roll off outside the
# field of view. Twofold oversampling halves the sample spacing
# :math:`\Delta k_x`, which doubles the field of view along the readout; the
# extent of k-space, and so the spatial resolution, is unchanged.
#
# The acquisition has 128 phase-encoding lines and 256 readout samples at
# half the phase-encoding spacing, in the grid units of a 128 matrix. The
# first trajectory component is :math:`k_x`, along the last image axis.

MATRIX = 128
phase_encodes = (torch.arange(MATRIX) - MATRIX // 2).float()
readout = (torch.arange(2 * MATRIX) - MATRIX).float() / 2
ky, kx = torch.meshgrid(phase_encodes, readout, indexing="ij")
trajectory = torch.stack([kx, ky, torch.zeros_like(kx)], dim=-1)

oversampled = bt.phantom(traj=trajectory).reshape(MATRIX, 2 * MATRIX)

# %%
#
# :func:`bartorch.remove_readout_oversampling` transforms the readout to the
# image domain, crops the field of view to the prescribed matrix and
# transforms back. Discarding the outer half of the readout samples instead
# keeps the doubled field of view and halves the resolution along the
# readout. The crop is compared with an acquisition without oversampling; the
# crop uses the unitary transform, whose normalization depends on the number
# of samples, so the two differ by the factor :math:`\sqrt{2}` between the
# transform lengths.

cropped = bartorch.remove_readout_oversampling(oversampled, MATRIX, axis=-1) / math.sqrt(2)
truncated = bartorch.resize(oversampled, (MATRIX, MATRIX)) / math.sqrt(2)
reference = bt.phantom(MATRIX, kspace=True)


def image_of(kspace):
    return bartorch.fft(kspace, axes=(-2, -1), inverse=True).abs()


def nrmse(estimate, target):
    return float((estimate - target).norm() / target.norm())


print(f"crop against no oversampling: NRMSE {nrmse(image_of(cropped), image_of(reference)):.1e}")

# %%

# sphinx_gallery_start_ignore
figure, axis = plt.subplots(figsize=(7.8, 4.1))
show(axis, image_of(oversampled), "256 readout samples", 0, 1.2)
axis.add_patch(Rectangle((63.5, -0.5), 128, 128, fill=False, edgecolor=BOX, linewidth=1.5))
plt.show()

difference = image_of(cropped) - image_of(reference)
figure, axes = plt.subplots(2, 2, figsize=(7.2, 7.0))
show(axes[0, 0], image_of(reference), "no oversampling", 0, 1.2)
show(axes[0, 1], image_of(cropped), "image-domain crop", 0, 1.2)
handle = show(axes[1, 0], difference, "crop - no oversampling", -0.01, 0.01, "RdBu_r")
figure.colorbar(handle, ax=axes[1, 0], shrink=0.8)
show(axes[1, 1], image_of(truncated), "k-space truncation", 0, 1.2)
plt.show()
# sphinx_gallery_end_ignore

# %%
#
# The image-domain crop reproduces the acquisition without oversampling to the
# part of the truncation ringing that extends past the prescribed field of
# view along the readout and is cropped with it. Discarding k-space samples
# instead widens the pixel to twice the prescribed size along :math:`x`: the
# object occupies half of the matrix.
#
# Apodization
# -----------
#
# An acquisition truncated at the edge of k-space is the object's spectrum
# multiplied by a rectangle, and its image the object convolved with a sinc.
# The first side lobe of the sinc is 22 % of its peak, which appears as an
# overshoot of about 9 % at a step edge and as ringing that decays over a few
# pixels. An apodization window rolls the data off towards the edge of
# k-space, which lowers the side lobes and widens the main lobe of the point
# spread function [#bernstein]_.
#
# The phantom is acquired on a 64 matrix and reconstructed on a 128 grid by
# zero-filling, which interpolates the image and makes the ringing visible
# between the pixels of the acquired grid.

GRID, ACQUIRED = 128, 64
image = bt.phantom(GRID).abs()
measured = bt.phantom(ACQUIRED, kspace=True)


def reconstruct(kspace):
    return bartorch.fft(bartorch.resize(kspace, (GRID, GRID)), axes=(-2, -1), inverse=True).abs()


# %%
#
# :func:`bartorch.fermi_window` sets the radius of the half height and the
# width of the transition separately, and keeps a wide passband.
# :func:`bartorch.hann_window` tapers from the centre of k-space to zero at
# the edge. The point spread functions are evaluated on a grid eight times
# finer than the acquired one: the side lobes of an unwindowed acquisition
# have their zeros at the pixels of the acquired grid, so a point spread
# function read off that grid shows none.

fermi = bartorch.fermi_window((ACQUIRED, ACQUIRED), radius=0.8, width=0.08)
hann = bartorch.hann_window((ACQUIRED, ACQUIRED))
windows = {"no window": torch.ones(ACQUIRED, ACQUIRED), "Fermi": fermi, "Hann": hann}

UPSAMPLE = 8
offset = (torch.arange(UPSAMPLE * ACQUIRED) - UPSAMPLE * ACQUIRED // 2) / UPSAMPLE


def psf(window):
    fine = bartorch.resize(window.to(torch.complex64), (UPSAMPLE * ACQUIRED,) * 2)
    profile = bartorch.fft(fine, axes=(-2, -1), inverse=True).abs()[UPSAMPLE * ACQUIRED // 2]
    return profile / profile.max()


def fwhm(profile):
    above = torch.nonzero(profile > 0.5).flatten()
    return float(offset[above[-1]] - offset[above[0]])


# %%
#
# :func:`bartorch.apodize` multiplies k-space by either window over the axes
# it names. The ringing is measured as the standard deviation of the image
# over the parenchyma within six pixels of the skull, where the object is
# uniform, and the resolution as the full width at half maximum of the point
# spread function, in pixels of the acquired grid.

reconstructions = {
    "no window": reconstruct(measured),
    "Fermi": reconstruct(bartorch.apodize(measured, kind="fermi", radius=0.8, width=0.08)),
    "Hann": reconstruct(bartorch.apodize(measured, kind="hann")),
}

skull = (image > 0.9).float()[None, None]
near_skull = torch.nn.functional.max_pool2d(skull, 13, stride=1, padding=6)[0, 0] > 0
flat = torch.nn.functional.max_pool2d(
    (image - 0.2).abs().gt(1e-3).float()[None, None], 5, stride=1, padding=2
)[0, 0].eq(0)
ringing_region = near_skull & flat

print(f"{'window':>9}  {'ringing':>7}  {'side lobe':>9}  {'FWHM':>7}")
metrics = {}
for name, window in windows.items():
    profile = psf(window)
    ringing = float(reconstructions[name][ringing_region].std())
    lobe = float(profile[offset.abs() > 3].max())
    metrics[name] = (ringing, fwhm(profile))
    print(f"{name:>9}  {ringing:7.4f}  {lobe:9.1e}  {fwhm(profile):5.2f} px")

# %%

# sphinx_gallery_start_ignore
ZOOM = (slice(36, 92), slice(4, 60))
figure, axes = plt.subplots(2, 2, figsize=(7.0, 7.2))
show(axes[0, 0], image[ZOOM], "object", 0.0, 0.5)
for axis, (name, estimate) in zip(axes.flat[1:], reconstructions.items(), strict=True):
    show(axis, estimate[ZOOM], f"{name}, FWHM {metrics[name][1]:.2f} px", 0.0, 0.5)
plt.show()

ROW = 64
figure, axes = plt.subplots(1, 2, figsize=(7.8, 3.4))
for name, window in windows.items():
    axes[0].semilogy(
        offset, psf(window).clamp_min(1e-4), color=COLOURS[name], linewidth=1.2, label=name
    )
axes[0].set_xlim(-8, 8)
axes[0].set_ylim(1e-4, 1.5)
axes[0].set_xlabel("offset [pixels, 64 matrix]")
axes[0].set_title("|PSF|")
axes[1].plot(image[ROW], color="#8a8a8a", linewidth=3.0, label="object")
for name, estimate in reconstructions.items():
    axes[1].plot(estimate[ROW], color=COLOURS[name], linewidth=1.2, label=name)
axes[1].legend(loc="upper right")
axes[1].set_xlim(0, 40)
axes[1].set_ylim(-0.05, 1.2)
axes[1].set_xlabel("pixel of the 128 grid")
axes[1].set_title(f"row {ROW}, across the skull")
plt.show()
# sphinx_gallery_end_ignore

# %%
#
# The zoomed panels show the left edge of the skull, displayed from zero to
# half the skull intensity. Without a window the ringing is visible as bands
# parallel to the skull across the adjacent parenchyma. The Fermi window lowers the far side lobes
# by a factor of three and the ringing next to the skull by about a fifth, and
# widens the point spread function by half a pixel; the Hann window removes
# the ringing almost entirely and doubles the full width at half maximum,
# which blurs the skull into the parenchyma. The Fermi window's wide passband
# is the usual compromise for anatomical imaging.
#
# Radial and separable windows
# ----------------------------
#
# The one-dimensional kernel is extended over k-space either on the
# Euclidean norm of the normalized coordinates (``geometry="radial"``, an
# ellipse) or as a product along each axis (``geometry="separable"``), which
# retains more of the corners of k-space. Bernstein et al. [#bernstein]_ give
# the ratio of the two at the diagonal Nyquist point as 52.4 % in two
# dimensions and 50.7 % in three, for a Fermi window of transition width
# 10/128. The radial window is 0.5 there; the separable one is the product of
# the one-dimensional kernel at :math:`1/\sqrt{d}` along each of :math:`d`
# axes. The one-dimensional kernel at :math:`u` is the window at the centre of
# a grid with its radius moved to :math:`1 - u`.

TRANSITION = 10.0 / 128


def kernel(u):
    return float(bartorch.fermi_window((4, 4), radius=1.0 - u, width=TRANSITION)[2, 2])


for dimensions in (2, 3):
    separable = kernel(1 / math.sqrt(dimensions)) ** dimensions
    print(f"{dimensions}D: radial / separable at the diagonal {100 * 0.5 / separable:.1f} %")

# %%

# sphinx_gallery_start_ignore
radial = bartorch.fermi_window((128, 128), width=TRANSITION, geometry="radial")
product = bartorch.fermi_window((128, 128), width=TRANSITION, geometry="separable")
figure, axes = plt.subplots(1, 2, figsize=(7.2, 3.4))
show(axes[0], radial, "radial", cmap="viridis")
handle = show(axes[1], product, "separable", cmap="viridis")
figure.colorbar(handle, ax=axes, shrink=0.9)
plt.show()

diagonal = torch.arange(64, 128)
radius = (diagonal - 64) / 64 * math.sqrt(2)
figure, axis = plt.subplots(figsize=(6.4, 3.2))
axis.plot(radius, radial[diagonal, diagonal], color=COLOURS["Fermi"], linewidth=1.4, label="radial")
axis.plot(
    radius, product[diagonal, diagonal], color=COLOURS["Hann"], linewidth=1.4, label="separable"
)
axis.axvline(1.0, color="#8a8a8a", linewidth=1.0, linestyle=":")
axis.set_xlabel("radius along the diagonal / Nyquist radius")
axis.set_ylabel("window")
axis.set_title("window along the diagonal")
axis.legend(loc="lower left")
plt.show()
# sphinx_gallery_end_ignore

# %%
#
# The radial window has the more isotropic point spread function and the
# higher signal-to-noise ratio; the separable one the better resolution along
# the diagonals. Both normalize each axis to its own Nyquist edge, so a
# rectangular matrix receives an ellipse matched to its grid.
#
# References
# ----------
#
# .. [#bernstein] Bernstein MA, Fain SB, Riederer SJ. Effect of windowing and
#    zero-filled reconstruction of MRI data on spatial resolution and
#    acquisition strategy. *J Magn Reson Imaging* 14(3):270-280 (2001).
#    https://doi.org/10.1002/jmri.1183
