Frequency-domain Light

class KSpaceLight(dim, kx, ky, wvl, field=None, source_pitch=None, source_dim=None, device='cpu')[source]

Complex optical field defined on a k-space (kx, ky) grid.

Unlike Light, which represents a complex field on a spatial (x, y) plane and carries spatial concepts such as pitch and bandwidth, KSpaceLight represents a field that already lives in k-space, i.e. on a grid of wave-vector coordinates (kx, ky) measured in radians per meter.

This is a lightweight, differentiable container. It does not perform any FFT internally, and it does not implement propagation or rendering logic. It only stores the complex k-space field and exposes complex-field accessors (shared in spirit with Light) together with k-space coordinate utilities useful for wide-angle propagation and angular rendering.

Parameters:
__init__(dim, kx, ky, wvl, field=None, source_pitch=None, source_dim=None, device='cpu')[source]

Create a k-space light container.

Parameters:
  • dim (tuple) – Field dimensions (B, Ch, R, C) for batch, channels, rows, cols.

  • kx (torch.Tensor) – 1D tensor of shape (C,), wave-vector coordinate along the x axis in radians per meter. Must be real floating-point, finite, strictly increasing and uniform (a singleton is valid).

  • ky (torch.Tensor) – 1D tensor of shape (R,), wave-vector coordinate along the y axis in radians per meter, with the same constraints as kx.

  • wvl (float or list) – Wavelength in meters. Either a single value or a list with length equal to the channel dimension.

  • field (torch.Tensor, optional) – Complex field [B, Ch, R, C]. If None, it is initialized to complex ones. Stored directly without detaching so that an existing computation graph is preserved.

  • source_pitch (float, optional) – Spatial pitch of the original spatial Light before Fourier conversion (metadata only).

  • source_dim (tuple, optional) – Original spatial Light dimensions (metadata only).

  • device (str) – Device for computation (‘cpu’, ‘cuda:0’, etc.).

Examples

>>> R, C = 256, 256
>>> kx = torch.linspace(-1e7, 1e7, C)
>>> ky = torch.linspace(-1e7, 1e7, R)
>>> klight = KSpaceLight(dim=(1, 1, R, C), kx=kx, ky=ky, wvl=633e-9)
_k0_scalar(c=None)[source]

Return one wavenumber for a coordinate grid, preserving tensor gradients.

Return type:

Union[float, Tensor]

Parameters:

c (int | None)

get_field(c=None)[source]

Return the complex field, optionally for a single channel.

Return type:

Tensor

Parameters:

c (int | None)

set_field(field, c=None)[source]

Set the complex field, optionally for a single channel.

Autograd is preserved: the incoming tensor is not detached.

Return type:

None

Parameters:
get_real(c=None)[source]

Return the real part of the field, optionally for a single channel.

Return type:

Tensor

Parameters:

c (int | None)

set_real(real, c=None)[source]

Set the real part of the field (keeps imaginary part). Autograd preserved.

Return type:

None

Parameters:
get_imag(c=None)[source]

Return the imaginary part of the field, optionally for a single channel.

Return type:

Tensor

Parameters:

c (int | None)

set_imag(imag, c=None)[source]

Set the imaginary part of the field (keeps real part). Autograd preserved.

Return type:

None

Parameters:
get_amplitude(c=None)[source]

Return the amplitude of the field, optionally for a single channel.

Return type:

Tensor

Parameters:

c (int | None)

set_amplitude(amplitude, c=None)[source]

Set the amplitude of the field (keeps phase). Autograd preserved.

Return type:

None

Parameters:
get_phase(c=None)[source]

Return the phase of the field (radians), optionally for a single channel.

Return type:

Tensor

Parameters:

c (int | None)

set_phase(phase, c=None)[source]

Set the phase of the field (keeps amplitude). Autograd preserved.

Return type:

None

Parameters:
clone()[source]

Create a deep copy. The field clone preserves the computation graph.

Return type:

KSpaceLight

shape()[source]

Return the shape of the field tensor.

Return type:

Size

get_channel()[source]

Return the number of channels.

Return type:

int

get_device()[source]

Return the device of the k-space light.

Return type:

str

to(device)[source]

Move field and coordinate tensors to the given device (in place).

Returns self to allow chaining. The field move preserves autograd.

Return type:

KSpaceLight

Parameters:

device (str)

get_kx_ky()[source]

Return the 1D kx and ky coordinate tensors.

Return type:

Tuple[Tensor, Tensor]

get_k_grid()[source]

Return 2D meshgrids (kx_grid, ky_grid), each of shape (R, C).

Return type:

Tuple[Tensor, Tensor]

get_k_limits()[source]

Return (kx_min, kx_max, ky_min, ky_max).

Return type:

Tuple[Tensor, Tensor, Tensor, Tensor]

get_k_spacing()[source]

Return (dkx, dky), assuming uniform spacing.

For a degenerate axis of length 1, the corresponding spacing is 0.

Return type:

Tuple[Tensor, Tensor]

get_k0(c=None)[source]

Return k0 = 2*pi/wavelength, preserving tensor dtype and gradients.

A numeric scalar returns a float. Numeric lists retain the historical float32 vector output. Tensor wavelengths, including scalar tensors or lists of tensors, retain their precision and graph on the field device.

Return type:

Union[float, Tensor]

Parameters:

c (int | None)

get_k_radius_grid()[source]

Return the sampled k-space radius sqrt(kx^2 + ky^2), shape (R, C).

Return type:

Tensor

get_valid_mask(c=None)[source]

Return the propagating-wave mask kx^2 + ky^2 <= k0^2, shape (R, C).

Return type:

Tensor

Parameters:

c (int | None)

get_direction_grid(c=None)[source]

Return unit direction components (dir_x, dir_y, dir_z, valid_mask).

dir_x = kx/k0, dir_y = ky/k0, dir_z = sqrt(max(1 - dir_x^2 - dir_y^2, 0)). valid_mask marks directions inside the unit circle (propagating waves).

Return type:

Tuple[Tensor, Tensor, Tensor, Tensor]

Parameters:

c (int | None)

get_theta_phi_grid(c=None)[source]

Return angular coordinates (theta, phi, valid_mask).

theta = asin(clamp(sqrt(dir_x^2 + dir_y^2), 0, 1)) is the polar angle from the optical axis; phi = atan2(dir_y, dir_x) is the azimuth.

Return type:

Tuple[Tensor, Tensor, Tensor]

Parameters:

c (int | None)

KSpaceLight.get_intensity(c=None)[source]

Return field.abs().square() for all channels, or for channel c.

Parameters:

c – Optional channel index.

Return type:

torch.Tensor