"""
===================================
EPI Nyquist ghost and ramp sampling
===================================

An echo-planar readout acquires k-space in a train of lines of alternating
readout gradient polarity, and two corrections are applied before the lines
form a Cartesian k-space. A timing error between the readout gradient and the
ADC, and eddy currents, displace the echoes of the reversed lines relative to
the forward ones; the resulting odd/even phase produces a Nyquist ghost, a
copy of the object displaced by half the field of view along the
phase-encoding direction. Sampling during the ramps of the readout gradient
shortens the echo spacing, and the samples, uniform in time, are not uniform
in :math:`k_x`.

This example simulates an eight-channel EPI acquisition of a Shepp-Logan
phantom with a known ADC delay and constant phase error, estimates the
odd/even phase from a three-line navigator, and compares the ghosted and the
corrected image with the delay-free image. It then resamples a ramp-sampled
readout onto a uniform :math:`k_x` grid and compares the result with linear
interpolation.

**Learning objectives**

* Relate an ADC delay to a linear phase in hybrid space and to the Nyquist
  ghost at half the field of view.
* Estimate the odd/even phase from a three-line navigator with
  :func:`~bartorch.tools.estimate_epi_phase`, and apply it to the reversed
  lines with :func:`~bartorch.tools.correct_lines`.
* Quantify the ghost by the ghost-to-signal ratio.
* Resample a ramp-sampled readout with
  :func:`~bartorch.tools.epi_ramp_operator`, and check the sampling condition
  under which the resampling is exact.
"""

# %%

# sphinx_gallery_start_ignore
import matplotlib.pyplot as plt

COLOURS = {"ghosted": "#e8a33d", "corrected": "#3dbde8", "linear": "#e8a33d"}


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 numpy as np
import torch

import bartorch
import bartorch.tools as bt

# %%
#
# The Nyquist ghost
# -----------------
#
# A delay :math:`\delta` of the ADC relative to the readout gradient shifts
# the echo of every line by the same time: towards positive :math:`k_x` on a
# forward line and towards negative :math:`k_x` on a reversed one. In hybrid
# space -- after the inverse Fourier transform along the readout -- a shift of
# :math:`\delta` dwell times is a linear phase :math:`\pi \delta u` over the
# readout coordinate :math:`u \in [-1, 1]`. A constant phase :math:`\phi_0`,
# from eddy currents or a :math:`B_0` offset during the readout, adds to it
# with the same alternating sign. The phase difference between odd and even
# lines modulates k-space at the Nyquist frequency of the phase-encoding
# direction, and the image acquires a ghost at half the field of view.
#
# Forward lines carry :math:`+(\pi \delta u + \phi_0)` and reversed lines
# :math:`-(\pi \delta u + \phi_0)`. A reversed line is stored in the order it
# was digitised, that is flipped along the readout.

SIZE = 128
DELAY = 0.2  # dwell times
PHASE_0 = 0.25  # rad
SLOPE = math.pi * DELAY * (SIZE - 1) / SIZE  # rad over u in [-1, 1]

kspace = bt.phantom(SIZE, kspace=True, coils=8)[:, 0]  # (coils, ky, kx)
hybrid = bartorch.fft(kspace, axes=(-1,), inverse=True, unitary=True)
u = torch.linspace(-1.0, 1.0, SIZE)


def acquire(row, polarity):
    """One readout of polarity +1 or -1, digitised in the order it was played."""
    delayed = row * torch.polar(torch.ones(SIZE), polarity * (SLOPE * u + PHASE_0))
    line = bartorch.fft(delayed, axes=(-1,), unitary=True)
    return torch.flip(line, [-1]) if polarity < 0 else line


train = [(acquire(hybrid[:, ky], 1 - 2 * (ky % 2)), ky % 2 == 1) for ky in range(SIZE)]

# %%
#
# The navigator
# -------------
#
# Three lines of alternating polarity acquired without phase-encoding blips
# sample the same line of k-space, so the phase between the reversed line and
# the mean of its two neighbours is the odd/even phase alone.
# :func:`~bartorch.tools.estimate_epi_phase` fits a polynomial to it,
# weighted by the signal magnitude and summed over coils. The reversed
# navigator line is passed already flipped into readout order.

centre = hybrid[:, SIZE // 2]
navigator = [acquire(centre, 1), torch.flip(acquire(centre, -1), [-1]), acquire(centre, 1)]

fit = bt.estimate_epi_phase(navigator)
print(f"fitted     constant {float(fit[0]):+.3f}  linear {float(fit[1]):+.3f} rad")
print(f"impressed  constant {2 * PHASE_0:+.3f}  linear {2 * SLOPE:+.3f} rad")

# %%
#
# The fitted phase is twice the impressed one, because the navigator measures
# the difference between a forward and a reversed line. The correction is
# applied to the reversed lines only: it brings them into phase with the
# forward lines, whose remaining phase is common to all lines and does not
# change the magnitude image.
#
# The ghost is measured by the ghost-to-signal ratio: the mean signal outside
# the object divided by the mean signal inside it.


def reconstruct(lines):
    coil_images = bartorch.fft(torch.stack(lines, dim=1), axes=(-2, -1), inverse=True, unitary=True)
    return bartorch.rss(coil_images, axes=(0,)).abs()


ideal = reconstruct([kspace[:, ky] for ky in range(SIZE)])
ghosted = reconstruct(bt.correct_lines(train))
corrected = reconstruct(bt.correct_lines(train, fit))

inside = ideal > 0.05 * ideal.max()
outside = ideal < 0.01 * ideal.max()


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


def ghost_to_signal(image):
    return float(image[outside].mean() / image[inside].mean())


for name, image in (("delay-free", ideal), ("flipped only", ghosted), ("corrected", corrected)):
    print(
        f"{name:>12}: ghost-to-signal {100 * ghost_to_signal(image):5.2f} %, "
        f"NRMSE {nrmse(image, ideal):.1e}"
    )

# %%

# sphinx_gallery_start_ignore
peak = float(ideal.max())
figure, axes = plt.subplots(1, 3, figsize=(7.8, 3.0))
for axis, (values, title) in zip(
    axes, ((ideal, "delay-free"), (ghosted, "flipped only"), (corrected, "corrected")), strict=True
):
    show(axis, values / peak, title, 0, 0.3)
plt.show()

readout = bartorch.fft(torch.stack(navigator), axes=(-1,), inverse=True, unitary=True)
cross = (0.5 * (readout[0] + readout[2]) * readout[1].conj()).sum(0)
keep = (cross.abs() > 0.1 * cross.abs().max()).numpy()
fitted = (fit[0] + fit[1] * u.double()).numpy()
figure, axes = plt.subplots(1, 2, figsize=(7.8, 3.4))
axes[0].plot(u[keep], torch.angle(cross)[keep], ".", color="#8a8a8a", label="navigator")
axes[0].plot(u, fitted, color=COLOURS["corrected"], label="first-order fit")
axes[0].set_xlabel("readout coordinate u")
axes[0].set_ylabel("odd/even phase [rad]")
axes[0].set_title("navigator phase")
axes[0].legend()
axes[1].semilogy(
    ideal[:, SIZE // 2] / peak + 1e-6, color="#8a8a8a", linewidth=3.0, label="delay-free"
)
axes[1].semilogy(
    ghosted[:, SIZE // 2] / peak + 1e-6, color=COLOURS["ghosted"], label="flipped only"
)
axes[1].semilogy(
    corrected[:, SIZE // 2] / peak + 1e-6, "--", color=COLOURS["corrected"], label="corrected"
)
axes[1].set_ylim(1e-3, 30)
axes[1].set_xlabel("phase-encoding index")
axes[1].set_title("central column")
axes[1].legend(loc="upper center", ncol=2)
plt.show()
# sphinx_gallery_end_ignore

# %%
#
# The images are displayed from zero to 30 % of the peak, with the
# phase-encoding direction vertical. Without the phase correction the ghost
# appears at the top and bottom of the field of view and overlaps the object
# where it wraps. With the first-order fit the ghost-to-signal ratio returns
# to that of the delay-free image, whose signal outside the object is the
# truncation ringing of the phantom, since the simulated phase error is
# exactly first order. On measured data, higher
# orders of the phase, and phase errors that differ between lines of the same
# polarity, leave a residual ghost.
#
# Ramp sampling
# -------------
#
# Sampling during the ramps of the trapezoidal readout gradient places the
# samples at :math:`k_x(t) = \gamma \int_0^t G_x(\tau)\, d\tau`, which is
# denser on the ramps than on the plateau. The readout is the Fourier
# transform of an object that spans ``support`` pixels, so samples at any
# positions determine it, provided no two neighbouring samples are further
# apart than one over the support, the Nyquist spacing of that object.
# :func:`~bartorch.tools.epi_ramp_operator` [#bruder]_ is the regularized
# least-squares inverse of the transform at the sampled positions followed by
# the transform at the uniform ones.
#
# The readout gradient is a trapezoid whose ramps each take 30 % of the ADC
# window, sampled with 160 samples, for a one-dimensional object of 64 pixels;
# positions are in cycles per pixel.

SUPPORT, SAMPLES = 64, 160
profile = torch.zeros(SUPPORT, dtype=torch.complex128)
profile[16:48] = torch.linspace(0.3, 1.0, 32) * torch.exp(1j * torch.linspace(0.0, 2.0, 32))
pixels = (torch.arange(SUPPORT) - SUPPORT // 2).double()

time = torch.linspace(0.0, 1.0, SAMPLES, dtype=torch.float64)
RAMP = 0.3
gradient = torch.clamp(torch.minimum(time / RAMP, (1 - time) / RAMP), 0, 1)
sampled_at = torch.cumsum(gradient, 0)
sampled_at = sampled_at - sampled_at.mean()
sampled_at = 0.5 * sampled_at / sampled_at.abs().max()
uniform_at = torch.arange(SAMPLES, dtype=torch.float64) / SAMPLES - 0.5


def encode(positions):
    return torch.exp(-2j * math.pi * torch.outer(positions, pixels)) @ profile


measured, truth = encode(sampled_at), encode(uniform_at)

operator = bt.epi_ramp_operator(sampled_at, uniform_at, SUPPORT)
resampled = (measured.to(torch.complex64)[None] @ operator.T)[0].to(torch.complex128)

# %%
#
# Linear interpolation between neighbouring samples is the comparison. Both
# are assessed on the image profile, the inverse transform of the uniform
# samples.

linear = torch.complex(
    torch.from_numpy(np.interp(uniform_at, sampled_at, measured.real)),
    torch.from_numpy(np.interp(uniform_at, sampled_at, measured.imag)),
)


def image_profile(samples):
    return torch.fft.fftshift(torch.fft.ifft(torch.fft.ifftshift(samples))).abs()


step = float(torch.diff(sampled_at).max()) * SUPPORT
print(f"largest step x support: {step:.2f}  (the samples determine the object below 1)")
for name, estimate in (("band-limited resampling", resampled), ("linear interpolation", linear)):
    print(f"{name:>24}: image NRMSE {nrmse(image_profile(estimate), image_profile(truth)):.1e}")

# %%

# sphinx_gallery_start_ignore
figure, axis = plt.subplots(figsize=(7.2, 3.2))
axis.plot(time, gradient, color="#8a8a8a", label="$G_x / G_{max}$")
axis.plot(time, 2 * sampled_at, color=COLOURS["corrected"], label="$k_x / k_{max}$, sampled")
axis.plot(time, 2 * uniform_at, "--", color=COLOURS["ghosted"], label="$k_x / k_{max}$, uniform")
axis.set_xlabel("time in the ADC window")
axis.set_title("trapezoidal readout gradient")
axis.legend(loc="center left", bbox_to_anchor=(1.0, 0.5))
plt.show()

figure, axes = plt.subplots(1, 2, figsize=(7.8, 3.4))
position = torch.arange(SAMPLES) - SAMPLES // 2
reference = image_profile(truth)
axes[0].plot(position, reference, color="#8a8a8a", linewidth=3.0, label="uniform")
axes[0].plot(position, image_profile(resampled), color=COLOURS["corrected"], label="resampled")
axes[0].plot(position, image_profile(linear), "--", color=COLOURS["linear"], label="linear")
axes[0].set_xlim(-50, 50)
axes[0].set_xlabel("pixel")
axes[0].set_title("image profile")
axes[0].legend(loc="upper left")
for estimate, name, style in ((resampled, "resampled", "-"), (linear, "linear", "--")):
    axes[1].semilogy(
        position,
        (image_profile(estimate) - reference).abs() / reference.max() + 1e-9,
        style,
        color=COLOURS[name if name == "linear" else "corrected"],
        label=name,
    )
axes[1].set_xlim(-80, 80)
axes[1].set_ylim(1e-8, 10)
axes[1].set_xlabel("pixel")
axes[1].set_title("|error| / peak")
axes[1].legend(loc="upper left", ncol=2)
plt.show()
# sphinx_gallery_end_ignore

# %%
#
# The band-limited resampling reproduces the profile of a uniformly sampled
# readout to the precision of the operator, which is returned in single
# precision. Linear interpolation errs most on the plateau, where the samples
# are furthest apart; in the image its error is spread over the whole field of
# view, inside and outside the object, at up to a few per cent of the peak.
# The ratio printed above is the condition to check on a measured trajectory:
# where a step exceeds one over the support, the readout is aliased and no
# resampling recovers it.
#
# References
# ----------
#
# .. [#bruder] Bruder H, Fischer H, Reinfelder HE, Schmitt F. Image
#    reconstruction for echo planar imaging with nonequidistant k-space
#    sampling. *Magn Reson Med* 23(2):311-323 (1992).
#    https://doi.org/10.1002/mrm.1910230211
