Skip to content

HybridLens

Combines a GeoLens with a diffractive optical element (DOE). HybridLens performs coherent ray tracing to the DOE plane, then Angular Spectrum Method (ASM) propagation to the sensor — a hybrid ray–wave model for refractive lenses with DOE or metasurface phase elements.

deeplens.HybridLens

HybridLens(
    filename=None,
    device=None,
    dtype=torch.float64,
    primary_wvln=DEFAULT_WAVE,
    wvln_rgb=WAVE_RGB,
    obj_depth=DEPTH,
)

Bases: Lens

Hybrid refractive-diffractive lens using a differentiable ray-wave model.

Combines a GeoLens (refractive module) with a diffractive optical element (DOE) placed behind it. The pipeline is: (1) coherent ray tracing through the embedded GeoLens to obtain a complex wavefront at the DOE plane (including all geometric aberrations); (2) DOE phase modulation applied to the wavefront; (3) Angular Spectrum Method (ASM) propagation from the DOE to the sensor plane to produce the final intensity PSF.

This enables end-to-end gradient flow from image-quality metrics back to both refractive surface parameters and the DOE phase profile. Operates in torch.float64 by default for numerical stability of the wave-propagation step.

Attributes:

Name Type Description
geolens GeoLens

Embedded refractive module. The DOE plane is appended to its surface list as a Plane placeholder.

doe Binary2 or Pixel2D or Fresnel or Zernike or Grating

Diffractive optical element behind the refractive group.

foclen float

Focal length [mm], copied from the embedded GeoLens.

Reference

Xinge Yang et al., "End-to-End Hybrid Refractive-Diffractive Lens Design with Differentiable Ray-Wave Model," SIGGRAPH Asia 2024.

Initialize a hybrid refractive-diffractive lens.

Parameters:

Name Type Description Default
filename str

Path to the lens configuration JSON file. Defaults to None.

None
device str

Computation device ('cpu' or 'cuda'). Defaults to None.

None
dtype dtype

Data type for computations. Defaults to torch.float64.

float64
primary_wvln float

Primary design wavelength [µm]. Used as fallback when a method is called without an explicit wvln. Defaults to DEFAULT_WAVE.

DEFAULT_WAVE
wvln_rgb list of float

Three wavelengths [µm] used for RGB computations, ordered [R, G, B]. Defaults to WAVE_RGB.

WAVE_RGB
obj_depth float

Default object depth [mm], used when a method is called without an explicit depth. Defaults to DEPTH.

DEPTH
Source code in deeplens-src/deeplens/hybridlens.py
def __init__(
    self,
    filename=None,
    device=None,
    dtype=torch.float64,
    primary_wvln=DEFAULT_WAVE,
    wvln_rgb=WAVE_RGB,
    obj_depth=DEPTH,
):
    """Initialize a hybrid refractive-diffractive lens.

    Args:
        filename (str, optional): Path to the lens configuration JSON file. Defaults to None.
        device (str, optional): Computation device ('cpu' or 'cuda'). Defaults to None.
        dtype (torch.dtype, optional): Data type for computations. Defaults to `torch.float64`.
        primary_wvln (float, optional): Primary design wavelength [µm].
            Used as fallback when a method is called without an explicit
            `wvln`. Defaults to `DEFAULT_WAVE`.
        wvln_rgb (list of float, optional): Three wavelengths [µm] used
            for RGB computations, ordered [R, G, B]. Defaults to `WAVE_RGB`.
        obj_depth (float, optional): Default object depth [mm], used
            when a method is called without an explicit `depth`. Defaults
            to `DEPTH`.
    """
    super().__init__(
        device=device,
        dtype=dtype,
        primary_wvln=primary_wvln,
        wvln_rgb=wvln_rgb,
        obj_depth=obj_depth,
    )

    # Load lens file
    if filename is not None:
        self.read_lens_json(filename)
    else:
        self.geolens = None
        self.doe = None
        # Set default sensor size and resolution if no file provided
        self.sensor_size = (8.0, 8.0)
        self.sensor_res = (2000, 2000)
        print(
            f"No lens file provided. Using default sensor_size: {self.sensor_size} mm, "
            f"sensor_res: {self.sensor_res} pixels. Use set_sensor() to change."
        )

    self.double()

read_lens_json

read_lens_json(filename)

Read the lens configuration from a JSON file.

Loads a GeoLens and associated DOE from the specified file. A Plane surface is appended to the GeoLens surface list as a placeholder for the DOE plane, matching the DOE aperture (square vs circular). Also sets self.foclen and the sensor size/resolution from the loaded GeoLens.

Supported DOE types: binary2, pixel2d, fresnel, zernike, grating.

Parameters:

Name Type Description Default
filename str

Path to the JSON configuration file. Must contain a "DOE" key with a "type" field.

required

Raises:

Type Description
ValueError

If the DOE type in the file is not supported.

Source code in deeplens-src/deeplens/hybridlens.py
def read_lens_json(self, filename):
    """Read the lens configuration from a JSON file.

    Loads a `GeoLens` and associated DOE from the specified file. A `Plane`
    surface is appended to the GeoLens surface list as a placeholder for the
    DOE plane, matching the DOE aperture (square vs circular). Also sets
    `self.foclen` and the sensor size/resolution from the loaded GeoLens.

    Supported DOE types: binary2, pixel2d, fresnel, zernike, grating.

    Args:
        filename (str): Path to the JSON configuration file. Must contain a
            "DOE" key with a "type" field.

    Raises:
        ValueError: If the DOE type in the file is not supported.
    """
    # Load geolens
    geolens = GeoLens(filename=filename, device=self.device)

    # Load DOE (diffractive surface)
    with open(filename, "r") as f:
        data = json.load(f)

        doe_dict = dict(data["DOE"])
        sensor_d = geolens.d_sensor.detach().clone()
        if "d_next" in doe_dict:
            doe_gap = torch.as_tensor(doe_dict["d_next"], device=self.device)
            doe_vertex_d = sensor_d - doe_gap
        elif "d" in doe_dict:
            # Backward compatibility for legacy hybrid files whose DOE
            # stored an absolute plane position.
            doe_vertex_d = torch.as_tensor(doe_dict["d"], device=self.device)
            doe_gap = sensor_d - doe_vertex_d
            doe_dict["d_next"] = float(doe_gap)
        else:
            raise ValueError("Hybrid DOE must define d_next.")

        doe_param_model = doe_dict["type"].lower()
        if doe_param_model == "binary2":
            doe = Binary2.init_from_dict(doe_dict)
        elif doe_param_model == "pixel2d":
            doe = Pixel2D.init_from_dict(doe_dict)
        elif doe_param_model == "fresnel":
            doe = Fresnel.init_from_dict(doe_dict)
        elif doe_param_model == "zernike":
            doe = Zernike.init_from_dict(doe_dict)
        elif doe_param_model == "grating":
            doe = Grating.init_from_dict(doe_dict)
        else:
            raise ValueError(f"Unsupported DOE parameter model: {doe_param_model}")
        self.doe = doe

    # Add a Plane/Phase surface to GeoLens (DOE placeholder).
    # Match the DOE's actual aperture (square vs circular) so that rays
    # outside the DOE region are correctly culled at the placeholder.
    # Split the existing last-surface-to-sensor gap at the DOE plane. Both
    # the geometric placeholder and wave-optics DOE use d_next.
    last_vertex_d = geolens.surf_d(-1).detach().clone()
    geolens.surfaces[-1].d_next = doe_vertex_d - last_vertex_d
    geolens.surfaces.append(
        Plane(
            d_next=doe.d_next,
            r=doe.r,
            mat2="air",
            is_square=doe.is_square,
        )
    )
    self.geolens = geolens
    self.foclen = geolens.foclen

    # Update hybrid lens sensor resolution and pixel size
    self.set_sensor(sensor_size=geolens.sensor_size, sensor_res=geolens.sensor_res)
    self.to(self.device)

write_lens_json

write_lens_json(lens_path)

Write the lens configuration to a JSON file.

Serialises the GeoLens surfaces (excluding the DOE placeholder) and the DOE configuration into a single JSON file that can be reloaded with read_lens_json.

Parameters:

Name Type Description Default
lens_path str

Output file path.

required
Source code in deeplens-src/deeplens/hybridlens.py
def write_lens_json(self, lens_path):
    """Write the lens configuration to a JSON file.

    Serialises the `GeoLens` surfaces (excluding the DOE placeholder) and the
    DOE configuration into a single JSON file that can be reloaded with
    `read_lens_json`.

    Args:
        lens_path (str): Output file path.
    """
    geolens = self.geolens
    data = {}
    data["info"] = geolens.lens_info if hasattr(geolens, "lens_info") else "None"
    data["foclen"] = round(geolens.foclen, 4)
    data["fnum"] = round(geolens.fnum, 4)
    data["r_sensor"] = round(geolens.r_sensor, 4)
    data["d_sensor"] = round(geolens.d_sensor.item(), 4)
    data["sensor_size"] = [round(i, 4) for i in geolens.sensor_size]
    data["sensor_res"] = geolens.sensor_res

    # Geolens
    data["surfaces"] = []
    for i, s in enumerate(geolens.surfaces[:-1]):
        surf_dict = s.surf_dict()

        # Exclude the DOE placeholder. Its sensor gap is folded back into
        # the final serialized refractive surface for reload compatibility.
        d_next = s.d_next
        if i == len(geolens.surfaces) - 2:
            d_next = d_next + geolens.surfaces[-1].d_next
        surf_dict["d_next"] = round(float(d_next), 3)

        data["surfaces"].append(surf_dict)

    # DOE
    data["DOE"] = self.doe.surf_dict()

    with open(lens_path, "w") as f:
        json.dump(data, f, indent=4)

analysis

analysis(save_name='./test.png')

Run a quick visual analysis of the hybrid lens.

Generates two figures: the 2D lens layout (saved to save_name) and the DOE phase map (saved to <save_name>_doe.png).

Parameters:

Name Type Description Default
save_name str

Base file path for the layout image. The DOE phase-map image path is formed by appending _doe.png. Defaults to "./test.png".

'./test.png'
Source code in deeplens-src/deeplens/hybridlens.py
def analysis(self, save_name="./test.png"):
    """Run a quick visual analysis of the hybrid lens.

    Generates two figures: the 2D lens layout (saved to `save_name`) and the
    DOE phase map (saved to `<save_name>_doe.png`).

    Args:
        save_name (str, optional): Base file path for the layout image. The
            DOE phase-map image path is formed by appending `_doe.png`.
            Defaults to "./test.png".
    """
    self.draw_layout(save_name=save_name)
    self.doe.draw_phase_map(save_name=f"{save_name}_doe.png")

double

double()

Convert the GeoLens and DOE to float64 precision.

Double precision is required for numerically stable phase accumulation during coherent ray tracing and ASM propagation. Called automatically by __init__.

Source code in deeplens-src/deeplens/hybridlens.py
def double(self):
    """Convert the GeoLens and DOE to `float64` precision.

    Double precision is required for numerically stable phase accumulation
    during coherent ray tracing and ASM propagation. Called automatically by
    `__init__`.
    """
    self.dtype = torch.float64
    self.geolens.astype(torch.float64)
    self.doe.astype(torch.float64)

refocus

refocus(foc_dist)

Refocus the hybrid lens to a given object distance.

Delegates to GeoLens.refocus, which adjusts the sensor distance; the DOE remains fixed relative to the refractive group (it is physically cemented to the lens barrel).

Parameters:

Name Type Description Default
foc_dist float

Target focus distance in [mm] (negative, towards the object).

required
Source code in deeplens-src/deeplens/hybridlens.py
def refocus(self, foc_dist):
    """Refocus the hybrid lens to a given object distance.

    Delegates to `GeoLens.refocus`, which adjusts the sensor distance; the
    DOE remains fixed relative to the refractive group (it is physically
    cemented to the lens barrel).

    Args:
        foc_dist (float): Target focus distance in [mm] (negative,
            towards the object).
    """
    self.geolens.refocus(foc_dist)

calc_scale

calc_scale(depth)

Calculate the object-to-image magnification scale factor.

Delegates to the embedded GeoLens.

Parameters:

Name Type Description Default
depth float

Object distance in [mm] (negative, towards the object).

required

Returns:

Name Type Description
scale float

Magnification factor (object height / image height), computed as \(-\text{depth} / \text{foclen}\).

Source code in deeplens-src/deeplens/hybridlens.py
def calc_scale(self, depth):
    """Calculate the object-to-image magnification scale factor.

    Delegates to the embedded `GeoLens`.

    Args:
        depth (float): Object distance in [mm] (negative, towards the
            object).

    Returns:
        scale (float): Magnification factor (object height / image height),
            computed as $-\\text{depth} / \\text{foclen}$.
    """
    return self.geolens.calc_scale(depth)

doe_field

doe_field(point, wvln=None, spp=SPP_COHERENT, upsample_factor=None)

Compute the complex wave field at the DOE plane via coherent ray tracing.

Similar to GeoLens.pupil_field, but evaluates the field at the last surface (DOE plane) instead of the exit pupil. The returned wavefront encodes amplitude, phase, and all diffraction-order information needed for subsequent DOE modulation and ASM propagation.

Parameters:

Name Type Description Default
point Tensor

Point source position, shape [3] or [1, 3] as [x, y, z]. x/y are in normalised sensor coordinates [-1, 1]; z is depth in [mm].

required
wvln float

Wavelength [µm]. When None (default), falls back to self.primary_wvln.

None
spp int

Number of rays to sample. Must be at least 1,000,000 for accurate coherent simulation. Defaults to SPP_COHERENT.

SPP_COHERENT
upsample_factor int or None

Field upsampling factor to meet the Nyquist sampling constraint. The field is sampled on a doe.res * upsample_factor grid with a doe.ps / upsample_factor pitch (same physical aperture, finer sampling). When None (default), a factor is chosen so the field resolution is close to 4000 x 4000.

None

Returns:

Name Type Description
wavefront Tensor

Complex wavefront at the DOE plane, shape [H, W] where H = W = doe.res[0] * upsample_factor.

psf_center list of float

Estimated PSF centre on the sensor in normalised coordinates [x, y].

Raises:

Type Description
AssertionError

If spp is less than 1,000,000 or the default dtype is not float64.

Source code in deeplens-src/deeplens/hybridlens.py
def doe_field(self, point, wvln=None, spp=SPP_COHERENT, upsample_factor=None):
    """Compute the complex wave field at the DOE plane via coherent ray tracing.

    Similar to `GeoLens.pupil_field`, but evaluates the field at the last
    surface (DOE plane) instead of the exit pupil. The returned wavefront
    encodes amplitude, phase, and all diffraction-order information needed
    for subsequent DOE modulation and ASM propagation.

    Args:
        point (torch.Tensor): Point source position, shape [3] or [1, 3] as
            [x, y, z]. x/y are in normalised sensor coordinates [-1, 1]; z is
            depth in [mm].
        wvln (float, optional): Wavelength [µm]. When None (default), falls
            back to `self.primary_wvln`.
        spp (int, optional): Number of rays to sample. Must be at least
            1,000,000 for accurate coherent simulation. Defaults to
            `SPP_COHERENT`.
        upsample_factor (int or None, optional): Field upsampling factor to
            meet the Nyquist sampling constraint. The field is sampled on a
            `doe.res * upsample_factor` grid with a `doe.ps / upsample_factor`
            pitch (same physical aperture, finer sampling). When None
            (default), a factor is chosen so the field resolution is close to
            4000 x 4000.

    Returns:
        wavefront (torch.Tensor): Complex wavefront at the DOE plane, shape
            [H, W] where H = W = `doe.res[0] * upsample_factor`.
        psf_center (list of float): Estimated PSF centre on the sensor in
            normalised coordinates [x, y].

    Raises:
        AssertionError: If `spp` is less than 1,000,000 or the default dtype
            is not `float64`.
    """
    wvln = self.primary_wvln if wvln is None else wvln
    assert spp >= 1_000_000, (
        "Coherent ray tracing spp is too small, "
        "which may lead to inaccurate simulation."
    )
    if self.dtype != torch.float64 or self.geolens.dtype != torch.float64:
        raise ValueError(
            "Coherent phase tracing requires float64 lens state; call double()."
        )

    geolens, doe = self.geolens, self.doe

    # Field-plane upsampling to satisfy the ASM Nyquist constraint
    if upsample_factor is None:
        upsample_factor = max(1, round(4000 / doe.res[0]))

    if point.dim() == 1:
        point = point.unsqueeze(0)
    point = point.to(self.device)

    # Calculate ray origin in the object space
    scale = geolens.calc_scale(point[:, 2].item())
    point_obj = point.clone()
    # sensor_size is (W, H): x scales with width [0], y with height [1].
    # (Matches the chief-ray center below and base Lens / DiffractiveLens.)
    point_obj[:, 0] = point[:, 0] * scale * geolens.sensor_size[0] / 2
    point_obj[:, 1] = point[:, 1] * scale * geolens.sensor_size[1] / 2

    # Determine ray center via chief ray
    pointc_chief_ray = geolens.psf_center(point_obj, method="chief_ray")[
        0
    ]  # shape [2]

    # Ray tracing to the DOE plane
    ray = geolens.sample_from_points(points=point_obj, num_rays=spp, wvln=wvln)
    ray.is_coherent = True
    ray, _ = geolens.trace(ray)
    ray = ray.prop_to(geolens.surf_d(-1))

    # Calculate full-resolution complex field for exit-pupil diffraction
    wavefront = forward_integral(
        ray.flip_xy(),
        ps=doe.ps / upsample_factor,
        ks=doe.res[0] * upsample_factor,
        pointc=torch.zeros_like(point[:, :2]),
    ).squeeze(0)  # shape [H, W]

    # Compute PSF center based on chief ray
    psf_center = [
        pointc_chief_ray[0] / geolens.sensor_size[0] * 2,
        pointc_chief_ray[1] / geolens.sensor_size[1] * 2,
    ]

    return wavefront, psf_center

psf

psf(points=None, wvln=None, ks=PSF_KS, **kwargs)

Compute a single-point monochromatic PSF using the ray-wave model.

The returned PSF includes all diffraction orders with physically correct diffraction efficiencies. The pipeline is: (1) coherent ray tracing through the GeoLens to obtain the complex wavefront at the DOE plane; (2) DOE phase modulation applied to the wavefront; (3) ASM propagation to the sensor, intensity calculation, cropping, and normalisation.

Parameters:

Name Type Description Default
points list or Tensor

[x, y, z] point source coordinates. x, y are in normalised sensor coordinates [-1, 1]; z is depth in [mm]. When None (default), uses [0.0, 0.0, -10000.0].

None
wvln float

Wavelength [µm]. When None (default), falls back to self.primary_wvln.

None
ks int or None

Output PSF patch size. When None, the centre half of the field is returned instead. Defaults to PSF_KS.

PSF_KS
**kwargs

Model-specific options. spp (int): number of coherent rays to sample, defaults to SPP_COHERENT. upsample_factor (int): field upsampling factor to meet the Nyquist sampling constraint; when None (default), a factor is chosen so the field resolution is close to 4000 x 4000.

{}

Returns:

Name Type Description
psf Tensor

Normalised PSF patch (sums to 1), shape [ks, ks] (or roughly half the field per side when ks is None). Returned in float32 precision.

Raises:

Type Description
ValueError

If the default dtype is not float64 (call double first).

Source code in deeplens-src/deeplens/hybridlens.py
def psf(self, points=None, wvln=None, ks=PSF_KS, **kwargs):
    """Compute a single-point monochromatic PSF using the ray-wave model.

    The returned PSF includes all diffraction orders with physically correct
    diffraction efficiencies. The pipeline is: (1) coherent ray tracing
    through the `GeoLens` to obtain the complex wavefront at the DOE plane;
    (2) DOE phase modulation applied to the wavefront; (3) ASM propagation to
    the sensor, intensity calculation, cropping, and normalisation.

    Args:
        points (list or torch.Tensor, optional): [x, y, z] point source
            coordinates. x, y are in normalised sensor coordinates [-1, 1];
            z is depth in [mm]. When None (default), uses
            [0.0, 0.0, -10000.0].
        wvln (float, optional): Wavelength [µm]. When None (default), falls
            back to `self.primary_wvln`.
        ks (int or None, optional): Output PSF patch size. When None, the
            centre half of the field is returned instead. Defaults to
            `PSF_KS`.
        **kwargs: Model-specific options. `spp` (int): number of coherent
            rays to sample, defaults to `SPP_COHERENT`. `upsample_factor`
            (int): field upsampling factor to meet the Nyquist sampling
            constraint; when None (default), a factor is chosen so the field
            resolution is close to 4000 x 4000.

    Returns:
        psf (torch.Tensor): Normalised PSF patch (sums to 1), shape [ks, ks]
            (or roughly half the field per side when `ks` is None). Returned
            in `float32` precision.

    Raises:
        ValueError: If the default dtype is not `float64` (call `double`
            first).
    """
    if points is None:
        points = [0.0, 0.0, -10000.0]
    spp = kwargs.get("spp", SPP_COHERENT)
    upsample_factor = kwargs.get("upsample_factor", None)
    wvln = self.primary_wvln if wvln is None else wvln
    # Check double precision
    if self.dtype != torch.float64 or self.geolens.dtype != torch.float64:
        raise ValueError(
            "Please call HybridLens.double() for accurate phase tracing."
        )

    # Check lens last surface
    assert isinstance(self.geolens.surfaces[-1], Phase) or isinstance(
        self.geolens.surfaces[-1], Plane
    ), "The last lens surface should be a DOE."
    geolens, doe = self.geolens, self.doe

    # Compute pupil field by coherent ray tracing
    if isinstance(points, list):
        point0 = torch.as_tensor(points, device=self.device, dtype=self.dtype)
    elif isinstance(points, torch.Tensor):
        point0 = points.to(device=self.device, dtype=self.dtype)
    else:
        raise ValueError("point should be a list or a torch.Tensor.")

    # Field-plane upsampling to satisfy the ASM Nyquist constraint
    if upsample_factor is None:
        upsample_factor = max(1, round(4000 / doe.res[0]))

    wavefront, psfc = self.doe_field(
        point=point0, wvln=wvln, spp=spp, upsample_factor=upsample_factor
    )
    wavefront = wavefront.squeeze(0)  # shape of [H, W]

    # DOE phase modulation. We have to flip the phase map because the
    # wavefront has been flipped. The phase map is upsampled (nearest, so
    # each flat DOE pixel is preserved) to the field resolution.
    phase_map = torch.flip(doe.get_phase_map(wvln), [-1, -2])
    if phase_map.shape != wavefront.shape:
        phase_map = F.interpolate(
            phase_map[None, None], size=wavefront.shape, mode="nearest"
        )[0, 0]
    wavefront = wavefront * torch.exp(1j * phase_map)

    # Propagate wave field to sensor plane
    h, w = wavefront.shape
    wavefront = F.pad(
        wavefront.unsqueeze(0).unsqueeze(0),
        [h // 2, h // 2, w // 2, w // 2],
        mode="constant",
        value=0,
    )
    sensor_field = AngularSpectrumMethod(
        wavefront,
        z=doe.d_next,
        wvln=wvln,
        ps=doe.ps / upsample_factor,
        padding=False,
    )

    # Compute PSF (intensity distribution)
    psf_inten = sensor_field.abs() ** 2
    psf_inten = (
        F.interpolate(
            psf_inten,
            scale_factor=geolens.sensor_res[0] / h,
            mode="bilinear",
            align_corners=False,
        )
        .squeeze(0)
        .squeeze(0)
    )

    # Calculate PSF center index and crop valid PSF region (Consider both interplation and padding)
    if ks is not None:
        h, w = psf_inten.shape[-2:]
        psfc_idx_i = ((2 - psfc[1]) * h / 4).round().long()
        psfc_idx_j = ((2 + psfc[0]) * w / 4).round().long()

        # Pad to avoid invalid edge region
        psf_inten_pad = F.pad(
            psf_inten,
            [ks // 2, ks // 2, ks // 2, ks // 2],
            mode="constant",
            value=0,
        )
        psf = psf_inten_pad[
            psfc_idx_i : psfc_idx_i + ks, psfc_idx_j : psfc_idx_j + ks
        ]
    else:
        h, w = psf_inten.shape[-2:]
        psf = psf_inten[
            int(h / 2 - h / 4) : int(h / 2 + h / 4),
            int(w / 2 - w / 4) : int(w / 2 + w / 4),
        ]

    # Normalize and convert to float precision.
    psf = psf / (psf.sum() + EPSILON)  # shape of [ks, ks] or [h, w]
    return diff_float(psf)

draw_layout

draw_layout(save_name='./DOELens.png', depth=-10000.0, ax=None, fig=None, dpi=600)

Draw the hybrid-lens layout with ray paths and wave-propagation arcs.

Renders the refractive elements via GeoLens.draw_lens_2d, traces rays at three field angles (on-axis, 0.707x, 0.99x full field), and overlays concentric arcs between the DOE and sensor to illustrate the wave-propagation region.

Parameters:

Name Type Description Default
save_name str

File path to save the figure (used only when ax is None). Defaults to "./DOELens.png".

'./DOELens.png'
depth float

Object depth [mm] for the traced rays. Defaults to -10000.0.

-10000.0
ax Axes

Pre-existing axes to draw into. If None, a new figure is created and saved.

None
fig Figure

Pre-existing figure. Required when ax is provided.

None
dpi int

Resolution used when saving a new figure. Defaults to 600.

600

Returns:

Name Type Description
ax Axes

The axes, returned only when ax was provided. When ax is None the figure is saved to save_name and nothing is returned.

fig Figure

The figure, returned only when ax was provided.

Source code in deeplens-src/deeplens/hybridlens.py
@torch.no_grad()
def draw_layout(
    self,
    save_name="./DOELens.png",
    depth=-10000.0,
    ax=None,
    fig=None,
    dpi=600,
):
    """Draw the hybrid-lens layout with ray paths and wave-propagation arcs.

    Renders the refractive elements via `GeoLens.draw_lens_2d`, traces rays
    at three field angles (on-axis, 0.707x, 0.99x full field), and overlays
    concentric arcs between the DOE and sensor to illustrate the
    wave-propagation region.

    Args:
        save_name (str, optional): File path to save the figure (used only
            when `ax` is None). Defaults to "./DOELens.png".
        depth (float, optional): Object depth [mm] for the traced rays.
            Defaults to -10000.0.
        ax (matplotlib.axes.Axes, optional): Pre-existing axes to draw into.
            If None, a new figure is created and saved.
        fig (matplotlib.figure.Figure, optional): Pre-existing figure.
            Required when `ax` is provided.
        dpi (int, optional): Resolution used when saving a new figure.
            Defaults to 600.

    Returns:
        ax (matplotlib.axes.Axes): The axes, returned only when `ax` was
            provided. When `ax` is None the figure is saved to `save_name`
            and nothing is returned.
        fig (matplotlib.figure.Figure): The figure, returned only when `ax`
            was provided.
    """
    geolens = self.geolens

    # Draw lens layout
    if ax is None:
        ax, fig = geolens.draw_lens_2d()
        save_fig = True
    else:
        save_fig = False

    # Draw DOE as orange Fresnel-style widget
    self.doe.draw_widget(ax, color="orange")

    # Draw light path
    color_list = ["#CC0000", "#006600", "#0066CC"]
    views = [
        0.0,
        float(np.rad2deg(geolens.rfov) * 0.707),
        float(np.rad2deg(geolens.rfov) * 0.99),
    ]
    arc_radi_list = [0.1, 0.4, 0.7, 1.0, 1.4, 1.8]
    num_rays = 11
    arc_half_angle = 20
    for i, view in enumerate(views):
        # Draw ray tracing
        ray = geolens.sample_point_source_2D(
            depth=depth,
            fov=view,
            num_rays=num_rays,
            entrance_pupil=True,
            wvln=self.wvln_rgb[2 - i],
        )
        ray.prop_to(-1.0)

        ray, ray_o_record = geolens.trace(ray=ray, record=True)
        ax, fig = geolens.draw_ray_2d(
            ray_o_record, ax=ax, fig=fig, color=color_list[i]
        )

        # Draw wave propagation
        # Calculate ray center for wave propagation visualization
        ray_center_doe = (
            ((ray.o * ray.is_valid.unsqueeze(-1)).sum(dim=0) / ray.is_valid.sum())
            .cpu()
            .numpy()
        )  # shape [3]
        ray.prop_to(geolens.d_sensor)  # shape [num_rays, 3]
        ray_center_sensor = (
            ((ray.o * ray.is_valid.unsqueeze(-1)).sum(dim=0) / ray.is_valid.sum())
            .cpu()
            .numpy()
        )  # shape [3]

        arc_radi = ray_center_sensor[2] - ray_center_doe[2]
        chief_theta = np.rad2deg(
            np.arctan2(
                ray_center_sensor[0] - ray_center_doe[0],
                ray_center_sensor[2] - ray_center_doe[2],
            )
        )
        theta1 = chief_theta - arc_half_angle
        theta2 = chief_theta + arc_half_angle

        for j in arc_radi_list:
            arc_radi_j = arc_radi * j
            arc = patches.Arc(
                (ray_center_sensor[2], ray_center_sensor[0]),
                arc_radi_j,
                arc_radi_j,
                angle=180.0,
                theta1=theta1,
                theta2=theta2,
                color=color_list[i],
            )
            ax.add_patch(arc)

    if save_fig:
        # Save figure
        ax.axis("off")
        ax.set_title("DOE Lens")
        fig.savefig(save_name, bbox_inches="tight", dpi=dpi)
        plt.close()
    else:
        return ax, fig

get_optimizer

get_optimizer(doe_lr=0.0001, lens_lr=[0.0001, 0.0001, 0.01, 1e-05])

Build an Adam optimiser for joint lens + DOE design.

Collects trainable parameters from both the GeoLens (surface thicknesses, curvatures, conic constants, aspheric coefficients) and the DOE phase profile into a single optimiser with per-group learning rates.

Parameters:

Name Type Description Default
doe_lr float

Learning rate for DOE phase parameters. Defaults to 1e-4.

0.0001
lens_lr list of float

Per-parameter-group learning rates for the GeoLens, ordered as [thickness_d, curvature_c, conic_k, aspheric_a]. Defaults to [1e-4, 1e-4, 1e-2, 1e-5].

[0.0001, 0.0001, 0.01, 1e-05]

Returns:

Name Type Description
optimizer Adam

Configured optimiser over all trainable parameters.

Source code in deeplens-src/deeplens/hybridlens.py
def get_optimizer(self, doe_lr=1e-4, lens_lr=[1e-4, 1e-4, 1e-2, 1e-5]):
    """Build an Adam optimiser for joint lens + DOE design.

    Collects trainable parameters from both the `GeoLens` (surface
    thicknesses, curvatures, conic constants, aspheric coefficients) and the
    DOE phase profile into a single optimiser with per-group learning rates.

    Args:
        doe_lr (float, optional): Learning rate for DOE phase parameters.
            Defaults to 1e-4.
        lens_lr (list of float, optional): Per-parameter-group learning rates
            for the GeoLens, ordered as
            [thickness_d, curvature_c, conic_k, aspheric_a]. Defaults to
            [1e-4, 1e-4, 1e-2, 1e-5].

    Returns:
        optimizer (torch.optim.Adam): Configured optimiser over all trainable
            parameters.
    """
    params = []
    params += self.geolens.get_optimizer_params(lrs=lens_lr)
    params += self.doe.get_optimizer_params(lr=doe_lr)

    optimizer = torch.optim.Adam(params)
    return optimizer

DOE Models

The diffractive element is configured by the DOE block in the lens JSON; its type field selects one of the phase parameterizations below. All subclass DiffractiveSurface, which defines the shared phase / propagation interface.

deeplens.diffractive_surface.DiffractiveSurface

DiffractiveSurface(
    d_next,
    res,
    fab_ps=0.001,
    fab_step=16,
    wvln0=0.55,
    mat="fused_silica",
    design_ps=None,
    is_square=True,
    device="cpu",
)

Bases: DeepObj

Base class for diffractive optical elements (DOEs).

A diffractive surface modulates the phase of an incident wave field; its optical behavior is simulated with wave optics. The phase profile is defined by phase_func in subclasses and converted into a wrapped, quantized phase map for the design wavelength. By default the DOE is designed for 0.55um, i.e. it has the highest 1st-order diffraction efficiency at 0.55um.

Attributes:

Name Type Description
d_next Tensor

Axial thickness to the next plane. [mm]

res tuple

DOE resolution as (H, W). [pixel]

ps float

Pixel size of the phase map (design pixel size if given, otherwise the fabrication pixel size). [mm]

w float

Physical width of the DOE. [mm]

h float

Physical height of the DOE. [mm]

is_square bool

Whether the aperture is treated as square.

r float

Aperture radius (half-diagonal / circumscribed-circle radius). [mm]

mat Material

DOE material.

wvln0 float

Design wavelength. [um]

n0 float

Refractive index of the material at wvln0.

fab_ps float

Fabrication pixel size. [mm]

fab_step int

Number of fabrication (quantization) levels.

x Tensor

x-coordinates of the grid. [H, W]. [mm]

y Tensor

y-coordinates of the grid. [H, W]. [mm]

Initialize a diffractive surface.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
res tuple or int

Resolution of the DOE as (H, W); an int is expanded to (res, res). [pixel]

required
fab_ps float

Fabrication pixel size. [mm]. Defaults to 0.001.

0.001
fab_step int

Number of fabrication (quantization) levels. Defaults to 16.

16
wvln0 float

Design wavelength. [um]. Defaults to 0.55.

0.55
mat str

Material name of the DOE. Defaults to "fused_silica".

'fused_silica'
design_ps float or None

Design pixel size; if None the fabrication pixel size is used as the phase-map pixel size. [mm]. Defaults to None.

None
is_square bool

Whether the aperture is square. Defaults to True.

True
device str

Device to place the DOE tensors on. Defaults to "cpu".

'cpu'
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def __init__(
    self,
    d_next,
    res,
    fab_ps=0.001,
    fab_step=16,
    wvln0=0.55,
    mat="fused_silica",
    design_ps=None,
    is_square=True,
    device="cpu",
):
    """Initialize a diffractive surface.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        res (tuple or int): Resolution of the DOE as (H, W); an int is
            expanded to (res, res). [pixel]
        fab_ps (float, optional): Fabrication pixel size. [mm]. Defaults to 0.001.
        fab_step (int, optional): Number of fabrication (quantization)
            levels. Defaults to 16.
        wvln0 (float, optional): Design wavelength. [um]. Defaults to 0.55.
        mat (str, optional): Material name of the DOE. Defaults to "fused_silica".
        design_ps (float or None, optional): Design pixel size; if None the
            fabrication pixel size is used as the phase-map pixel size. [mm].
            Defaults to None.
        is_square (bool, optional): Whether the aperture is square. Defaults to True.
        device (str, optional): Device to place the DOE tensors on. Defaults to "cpu".
    """
    # Geometry
    self.d_next = (
        d_next.detach().clone()
        if torch.is_tensor(d_next)
        else torch.tensor(d_next, dtype=torch.get_default_dtype())
    )
    if not self.d_next.is_floating_point():
        self.d_next = self.d_next.to(torch.get_default_dtype())
    self.res = (res, res) if isinstance(res, int) else res
    self.ps = fab_ps if design_ps is None else design_ps
    self.w = self.res[0] * self.ps
    self.h = self.res[1] * self.ps
    self.is_square = is_square
    # Surface radius: half-diagonal (circumscribed-circle radius) so it
    # is consistent with Phase / Surface conventions for square apertures.
    self.r = float(np.sqrt(self.w**2 + self.h**2) / 2)

    # Phase map
    self.mat = Material(mat)
    self.wvln0 = wvln0  # [um], design wavelength. Sometimes the maximum working wavelength is preferred.
    self.n0 = self.mat.refractive_index(
        self.wvln0
    )  # refractive index at design wavelength

    # Fabrication for DOE
    self.fab_ps = fab_ps  # [mm], fabrication pixel size
    self.fab_step = fab_step

    # x, y coordinates
    self.x, self.y = torch.meshgrid(
        torch.linspace(-self.w / 2, self.w / 2, self.res[1]),
        torch.linspace(self.h / 2, -self.h / 2, self.res[0]),
        indexing="xy",
    )

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a DOE from a dict.

Parameters:

Name Type Description Default
doe_dict dict

Dictionary of DOE parameters.

required

Returns:

Name Type Description
doe DiffractiveSurface

The constructed DOE instance.

Raises:

Type Description
NotImplementedError

Must be implemented by subclasses.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a DOE from a dict.

    Args:
        doe_dict (dict): Dictionary of DOE parameters.

    Returns:
        doe (DiffractiveSurface): The constructed DOE instance.

    Raises:
        NotImplementedError: Must be implemented by subclasses.
    """
    raise NotImplementedError

phase_func

phase_func()

Compute the raw phase profile (no wrapping, no quantization) at the design wavelength.

Returns:

Name Type Description
phase Tensor

Raw, unwrapped phase profile at the design wavelength. [H, W]. [rad]

Raises:

Type Description
NotImplementedError

Must be implemented by subclasses.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def phase_func(self):
    """Compute the raw phase profile (no wrapping, no quantization) at the design wavelength.

    Returns:
        phase (torch.Tensor): Raw, unwrapped phase profile at the design
            wavelength. [H, W]. [rad]

    Raises:
        NotImplementedError: Must be implemented by subclasses.
    """
    raise NotImplementedError

get_phase_map0

get_phase_map0()

Compute the phase map at the design wavelength with wrapping and quantization.

The raw phase from phase_func is wrapped into \([0, 2\pi)\) and then quantized to fab_step levels. The wrapped phase is equivalent to a height map whose maximum height corresponds to \(2\pi\) at the design wavelength.

Returns:

Name Type Description
phase0 Tensor

Wrapped, quantized phase map at the design wavelength. [H, W], range \([0, 2\pi)\). [rad]

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def get_phase_map0(self):
    """Compute the phase map at the design wavelength with wrapping and quantization.

    The raw phase from `phase_func` is wrapped into $[0, 2\\pi)$ and then
    quantized to `fab_step` levels. The wrapped phase is equivalent to a
    height map whose maximum height corresponds to $2\\pi$ at the design
    wavelength.

    Returns:
        phase0 (torch.Tensor): Wrapped, quantized phase map at the design
            wavelength. [H, W], range $[0, 2\\pi)$. [rad]
    """
    # Raw phase map at design wavelength
    phase0 = self.phase_func()

    # Phase wrapping and quantization
    phase0 = torch.remainder(phase0, 2 * torch.pi)
    phase0 = diff_quantize(phase0, levels=self.fab_step)
    return phase0

get_phase_map

get_phase_map(wvln)

Compute the phase map at the given wavelength.

The phase map is first computed at the design wavelength, then scaled to the requested wavelength accounting for the wavelength ratio and the material dispersion \((n - 1) / (n_0 - 1)\), and finally resampled (nearest) to the DOE resolution if needed.

Parameters:

Name Type Description Default
wvln float

Wavelength. [um]

required

Returns:

Name Type Description
phase_map Tensor

Phase map at the given wavelength. [H, W]. [rad]

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def get_phase_map(self, wvln):
    """Compute the phase map at the given wavelength.

    The phase map is first computed at the design wavelength, then scaled to
    the requested wavelength accounting for the wavelength ratio and the
    material dispersion $(n - 1) / (n_0 - 1)$, and finally resampled
    (nearest) to the DOE resolution if needed.

    Args:
        wvln (float): Wavelength. [um]

    Returns:
        phase_map (torch.Tensor): Phase map at the given wavelength.
            [H, W]. [rad]
    """
    # Phase map at design wavelength
    phase_map0 = self.get_phase_map0()

    # Phase map at given wavelength (implicitly converted to height map)
    n = self.mat.refractive_index(wvln)
    phase_map = phase_map0 * (self.wvln0 / wvln) * (n - 1) / (self.n0 - 1)

    # Interpolate to the desired resolution (skip if already matching)
    if phase_map.shape[-2:] != (self.res[0], self.res[1]):
        phase_map = (
            F.interpolate(
                phase_map.unsqueeze(0).unsqueeze(0), size=self.res, mode="nearest"
            )
            .squeeze(0)
            .squeeze(0)
        )

    return phase_map

forward

forward(wave)

Apply phase modulation, then propagate by this surface's d_next.

The input wave field may have a different pixel size and physical extent than the DOE; the phase map is resampled (nearest) to match the wave pixel size, then center-cropped or zero-padded to match the wave resolution before being applied as \(u \cdot e^{i\phi}\).

Parameters:

Name Type Description Default
wave ComplexWave

Input complex wave field, with field u of shape [B, 1, H, W].

required

Returns:

Name Type Description
wave ComplexWave

Output complex wave field after propagation and phase modulation, field u of shape [B, 1, H, W].

Reference

[1] https://github.com/vsitzmann/deepoptics function phaseshifts_from_height_map

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def forward(self, wave):
    """Apply phase modulation, then propagate by this surface's `d_next`.

    The input wave field may have a different pixel size and physical extent
    than the DOE; the phase map is resampled (nearest) to match the wave
    pixel size, then center-cropped or zero-padded to match the wave
    resolution before being applied as $u \\cdot e^{i\\phi}$.

    Args:
        wave (ComplexWave): Input complex wave field, with field `u` of
            shape [B, 1, H, W].

    Returns:
        wave (ComplexWave): Output complex wave field after propagation and
            phase modulation, field `u` of shape [B, 1, H, W].

    Reference:
        [1] https://github.com/vsitzmann/deepoptics function phaseshifts_from_height_map
    """
    # Compute phase map at the wave field wavelength, shape of [H, W]
    phase_map = self.get_phase_map(wave.wvln)

    # Consider the different pixel size between the wave field and the DOE
    if self.ps != wave.ps:
        scale = self.ps / wave.ps
        phase_map = (
            F.interpolate(
                phase_map.unsqueeze(0).unsqueeze(0),
                scale_factor=(scale, scale),
                mode="nearest",
            )
            .squeeze(0)
            .squeeze(0)
        )

    # Check if the field and phase map resolution (physical size) are the same
    wave_h, wave_w = wave.u.shape[-2:]
    phase_h, phase_w = phase_map.shape[-2:]
    if phase_h > wave_h or phase_w > wave_w:
        start_h = (phase_h - wave_h) // 2
        start_w = (phase_w - wave_w) // 2
        phase_map = phase_map[
            ..., start_h : start_h + wave_h, start_w : start_w + wave_w
        ]
    elif phase_h < wave_h or phase_w < wave_w:
        pad_top = (wave_h - phase_h) // 2
        pad_bottom = wave_h - phase_h - pad_top
        pad_left = (wave_w - phase_w) // 2
        pad_right = wave_w - phase_w - pad_left
        phase_map = F.pad(
            phase_map,
            (pad_left, pad_right, pad_top, pad_bottom),
            mode="constant",
            value=0,
        )

    wave.u = wave.u * torch.exp(1j * phase_map)
    wave.prop(self.d_next)
    return wave

__call__

__call__(wave)

Apply the DOE to a wave field (alias for forward).

Parameters:

Name Type Description Default
wave ComplexWave

Input complex wave field.

required

Returns:

Name Type Description
wave ComplexWave

Output complex wave field.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def __call__(self, wave):
    """Apply the DOE to a wave field (alias for `forward`).

    Args:
        wave (ComplexWave): Input complex wave field.

    Returns:
        wave (ComplexWave): Output complex wave field.
    """
    return self.forward(wave)

quantize_phase_map

quantize_phase_map(bits=16)

Quantize the design-wavelength phase map to a given number of levels.

Parameters:

Name Type Description Default
bits int

Number of quantization levels. Defaults to 16.

16

Returns:

Name Type Description
pmap_q Tensor

Quantized phase map. [H, W], range \([0, 2\pi)\). [rad]

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def quantize_phase_map(self, bits=16):
    """Quantize the design-wavelength phase map to a given number of levels.

    Args:
        bits (int, optional): Number of quantization levels. Defaults to 16.

    Returns:
        pmap_q (torch.Tensor): Quantized phase map. [H, W], range
            $[0, 2\\pi)$. [rad]
    """
    pmap = self.get_phase_map0()
    pmap_q = torch.round(pmap / (2 * torch.pi / bits)) * (2 * torch.pi / bits)
    return pmap_q

export_fab_phase_map

export_fab_phase_map(bits=16, save_path=None)

Generate a fabrication-resolution quantized phase map and save a checkpoint.

The phase map is upsampled from the design pixel size to the fabrication pixel size (bilinear) and quantized to bits levels. The DOE checkpoint is saved to save_path; the DOE object itself is left unchanged.

Parameters:

Name Type Description Default
bits int

Number of quantization levels. Defaults to 16.

16
save_path str or None

Checkpoint save path; if None a name encoding the fabrication resolution, pixel size, and bit depth is generated. Defaults to None.

None

Returns:

Name Type Description
pmap_q Tensor

Fabrication-resolution quantized phase map. [H_fab, W_fab], range \([0, 2\pi)\). [rad]

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def export_fab_phase_map(self, bits=16, save_path=None):
    """Generate a fabrication-resolution quantized phase map and save a checkpoint.

    The phase map is upsampled from the design pixel size to the fabrication
    pixel size (bilinear) and quantized to `bits` levels. The DOE checkpoint
    is saved to `save_path`; the DOE object itself is left unchanged.

    Args:
        bits (int, optional): Number of quantization levels. Defaults to 16.
        save_path (str or None, optional): Checkpoint save path; if None a
            name encoding the fabrication resolution, pixel size, and bit
            depth is generated. Defaults to None.

    Returns:
        pmap_q (torch.Tensor): Fabrication-resolution quantized phase map.
            [H_fab, W_fab], range $[0, 2\\pi)$. [rad]
    """
    # Fab resolution quantized pmap
    pmap = self.get_phase_map0()
    fab_res = int(self.ps / self.fab_ps * self.res[0])
    pmap = (
        F.interpolate(
            pmap.unsqueeze(0).unsqueeze(0),
            scale_factor=self.ps / self.fab_ps,
            mode="bilinear",
            align_corners=True,
        )
        .squeeze(0)
        .squeeze(0)
    )
    pmap_q = torch.round(pmap / (2 * torch.pi / bits)) * (2 * torch.pi / bits)

    # Save phase map
    if save_path is None:
        save_path = f"./doe_fab_{fab_res}x{fab_res}_{int(self.fab_ps * 1000)}um_{bits}bit.pth"
    self.save_ckpt(save_path=save_path)

    return pmap_q

activate_grad

activate_grad(activate=True)

Enable or disable gradients on the phase-map parameters.

Parameters:

Name Type Description Default
activate bool

Whether to require gradients. Defaults to True.

True

Raises:

Type Description
NotImplementedError

Must be implemented by subclasses.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def activate_grad(self, activate=True):
    """Enable or disable gradients on the phase-map parameters.

    Args:
        activate (bool, optional): Whether to require gradients. Defaults to True.

    Raises:
        NotImplementedError: Must be implemented by subclasses.
    """
    raise NotImplementedError

get_optimizer_params

get_optimizer_params(lr=None)

Build optimizer parameter groups for the phase-map parameters.

Parameters:

Name Type Description Default
lr float or None

Learning rate. Defaults to None.

None

Returns:

Name Type Description
params list

List of parameter group dicts for an optimizer.

Raises:

Type Description
NotImplementedError

Must be implemented by subclasses.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def get_optimizer_params(self, lr=None):
    """Build optimizer parameter groups for the phase-map parameters.

    Args:
        lr (float or None, optional): Learning rate. Defaults to None.

    Returns:
        params (list): List of parameter group dicts for an optimizer.

    Raises:
        NotImplementedError: Must be implemented by subclasses.
    """
    raise NotImplementedError

get_optimizer

get_optimizer(lr=None)

Create an Adam optimizer for the DOE phase-map parameters.

Parameters:

Name Type Description Default
lr float or None

Learning rate passed to get_optimizer_params. Defaults to None.

None

Returns:

Name Type Description
optimizer Adam

Optimizer over the DOE phase-map parameters.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def get_optimizer(self, lr=None):
    """Create an Adam optimizer for the DOE phase-map parameters.

    Args:
        lr (float or None, optional): Learning rate passed to
            `get_optimizer_params`. Defaults to None.

    Returns:
        optimizer (torch.optim.Adam): Optimizer over the DOE phase-map
            parameters.
    """
    params = self.get_optimizer_params(lr)
    optimizer = torch.optim.Adam(params)

    return optimizer

loss_quantization

loss_quantization(bits=16)

Compute the mean phase quantization error of the DOE.

Returns the mean absolute difference between the continuous phase map and its quantization to bits levels, used as a quantization-aware regularization loss.

Parameters:

Name Type Description Default
bits int

Number of quantization levels. Defaults to 16.

16

Returns:

Name Type Description
loss Tensor

Scalar mean absolute quantization error. [rad]

Reference

Quantization-aware Deep Optics for Diffractive Snapshot Hyperspectral Imaging.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def loss_quantization(self, bits=16):
    """Compute the mean phase quantization error of the DOE.

    Returns the mean absolute difference between the continuous phase map
    and its quantization to `bits` levels, used as a quantization-aware
    regularization loss.

    Args:
        bits (int, optional): Number of quantization levels. Defaults to 16.

    Returns:
        loss (torch.Tensor): Scalar mean absolute quantization error. [rad]

    Reference:
        Quantization-aware Deep Optics for Diffractive Snapshot Hyperspectral Imaging.
    """
    pmap = self.get_phase_map0()
    step = 2 * torch.pi / bits
    pmap_q = torch.round(pmap / step) * step
    loss = torch.mean(torch.abs(pmap - pmap_q))
    return loss

draw_phase_map

draw_phase_map(bits=None, save_name='./DOE_phase_map.png')

Save the design-wavelength phase map as a normalized image.

Parameters:

Name Type Description Default
bits int or None

Number of quantization levels; if given the phase map is quantized first, otherwise the continuous map is used. Defaults to None.

None
save_name str

Path to save the image. Defaults to "./DOE_phase_map.png".

'./DOE_phase_map.png'
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def draw_phase_map(self, bits=None, save_name="./DOE_phase_map.png"):
    """Save the design-wavelength phase map as a normalized image.

    Args:
        bits (int or None, optional): Number of quantization levels; if
            given the phase map is quantized first, otherwise the
            continuous map is used. Defaults to None.
        save_name (str, optional): Path to save the image. Defaults to
            "./DOE_phase_map.png".
    """
    if bits is not None:
        pmap = self.quantize_phase_map(bits)
    else:
        pmap = self.get_phase_map0()
    save_image(pmap, save_name, normalize=True)

draw_phase_map3d

draw_phase_map3d(bits=None, save_name='./DOE_phase_map3d.png')

Save a 3D scatter plot of the design-wavelength phase map.

Parameters:

Name Type Description Default
bits int or None

Number of quantization levels; if given the phase map is quantized first, otherwise the continuous map is used. Defaults to None.

None
save_name str

Path to save the image. Defaults to "./DOE_phase_map3d.png".

'./DOE_phase_map3d.png'
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def draw_phase_map3d(self, bits=None, save_name="./DOE_phase_map3d.png"):
    """Save a 3D scatter plot of the design-wavelength phase map.

    Args:
        bits (int or None, optional): Number of quantization levels; if
            given the phase map is quantized first, otherwise the
            continuous map is used. Defaults to None.
        save_name (str, optional): Path to save the image. Defaults to
            "./DOE_phase_map3d.png".
    """
    if bits is not None:
        pmap = self.quantize_phase_map(bits)
    else:
        pmap = self.get_phase_map0()

    pmap = pmap / 20.0
    x = np.linspace(-self.w / 2, self.w / 2, self.res[0])
    y = np.linspace(-self.h / 2, self.h / 2, self.res[1])
    X, Y = np.meshgrid(x, y)

    fig = plt.figure(figsize=(5, 5))
    ax = fig.add_subplot(111, projection="3d")
    ax.scatter(
        X.flatten(),
        Y.flatten(),
        pmap.cpu().numpy().flatten(),
        marker=".",
        s=0.01,
        c=pmap.cpu().numpy().flatten(),
        cmap="viridis",
    )
    ax.set_aspect("equal")
    ax.axis("off")
    fig.savefig(save_name, dpi=600, bbox_inches="tight")
    plt.close(fig)

draw_phase_map_fab

draw_phase_map_fab(save_name='./DOE_phase_map.png')

Save side-by-side images of the continuous and 16-level quantized phase maps.

Parameters:

Name Type Description Default
save_name str

Path to save the figure. Defaults to "./DOE_phase_map.png".

'./DOE_phase_map.png'
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def draw_phase_map_fab(self, save_name="./DOE_phase_map.png"):
    """Save side-by-side images of the continuous and 16-level quantized phase maps.

    Args:
        save_name (str, optional): Path to save the figure. Defaults to
            "./DOE_phase_map.png".
    """
    pmap = self.get_phase_map0()
    step = 2 * torch.pi / 16
    pmap_q = torch.round(pmap / step) * step

    fig, ax = plt.subplots(1, 2, figsize=(10, 5))
    ax[0].imshow(pmap.cpu().numpy(), vmin=0, vmax=2 * float(np.pi))
    ax[0].set_title(f"Phase map ({self.wvln0}um)", fontsize=10)
    ax[0].grid(False)
    fig.colorbar(ax[0].get_images()[0])

    ax[1].imshow(pmap_q.cpu().numpy(), vmin=0, vmax=2 * float(np.pi))
    ax[1].set_title(f"Quantized phase map ({self.wvln0}um)", fontsize=10)
    ax[1].grid(False)
    fig.colorbar(ax[1].get_images()[0])

    fig.savefig(save_name, dpi=600, bbox_inches="tight")
    plt.close(fig)

draw_cross_section

draw_cross_section(save_name='./DOE_cross_section.png')

Save a plot of the phase map along its main diagonal.

Parameters:

Name Type Description Default
save_name str

Path to save the figure. Defaults to "./DOE_cross_section.png".

'./DOE_cross_section.png'
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def draw_cross_section(self, save_name="./DOE_cross_section.png"):
    """Save a plot of the phase map along its main diagonal.

    Args:
        save_name (str, optional): Path to save the figure. Defaults to
            "./DOE_cross_section.png".
    """
    pmap = self.get_phase_map0()
    pmap = torch.diag(pmap).cpu().numpy()
    r = np.linspace(
        -self.w / 2 * float(np.sqrt(2)), self.w / 2 * float(np.sqrt(2)), self.res[0]
    )

    fig, ax = plt.subplots()
    ax.plot(r, pmap)
    ax.set_title(f"Phase map ({self.wvln0}um) cross section")
    fig.savefig(save_name, dpi=600, bbox_inches="tight")
    plt.close(fig)

draw_widget

draw_widget(ax, color='orange', linestyle='-', d=0.0)

Draw a 2D Fresnel-style cross-section of the DOE in a layout plot.

Plots the cross-section along the x-axis at y=0. For a square aperture the half-extent is the half-side (w/2); for a circular aperture it is the full radius r (= half-diagonal).

Parameters:

Name Type Description Default
ax Axes

Axes to draw on.

required
color str

Line color. Defaults to "orange".

'orange'
linestyle str

Line style. Defaults to "-".

'-'
d float

Derived global vertex position [mm]. Defaults to 0.

0.0
Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def draw_widget(self, ax, color="orange", linestyle="-", d=0.0):
    """Draw a 2D Fresnel-style cross-section of the DOE in a layout plot.

    Plots the cross-section along the x-axis at y=0. For a square aperture
    the half-extent is the half-side (`w/2`); for a circular aperture it is
    the full radius `r` (= half-diagonal).

    Args:
        ax (matplotlib.axes.Axes): Axes to draw on.
        color (str, optional): Line color. Defaults to "orange".
        linestyle (str, optional): Line style. Defaults to "-".
        d (float, optional): Derived global vertex position [mm].
            Defaults to 0.
    """
    d = float(d)
    max_offset = self.r / 100
    roc = self.r * 2
    x_half = self.w / 2 if self.is_square else self.r
    x = np.linspace(-x_half, x_half, 256)
    sag = roc * (1 - np.sqrt(1 - x**2 / roc**2))
    sag = max_offset - np.fmod(sag, max_offset)
    ax.plot(d + sag, x, color=color, linestyle=linestyle, linewidth=0.75)

surf_dict

surf_dict()

Serialize the DOE surface parameters into a dict.

Returns:

Name Type Description
surf_dict dict

Surface parameters (type, size, thickness, design wavelength, resolution, fabrication pixel size, and aperture shape flag) suitable for saving or reconstruction.

Source code in deeplens-src/deeplens/diffractive_surface/base_diffractive.py
def surf_dict(self):
    """Serialize the DOE surface parameters into a dict.

    Returns:
        surf_dict (dict): Surface parameters (type, size, thickness,
            design wavelength, resolution, fabrication pixel size, and
            aperture shape flag) suitable for saving or reconstruction.
    """
    surf_dict = {
        "type": self.__class__.__name__,
        "(size)": [round(self.w, 4), round(self.h, 4)],
        "d_next": round(self.d_next.item(), 4),
        "wvln0": round(self.wvln0, 4),
        "res": self.res,
        "fab_ps": self.fab_ps,
        "is_square": self.is_square,
    }

    return surf_dict

Polynomial (Binary-2) rotationally-symmetric phase profile.

deeplens.diffractive_surface.Binary2

Binary2(
    d_next,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    is_square=True,
    device="cpu",
)

Bases: DiffractiveSurface

Binary2 (Zemax-style) rotationally symmetric DOE surface.

Parameterizes the design-wavelength phase as an even polynomial in the radial coordinate, \(\phi(r) = \pi \sum_{i=1}^{5} \alpha_{2i}\, r^{2i}\), with coefficients alpha2, alpha4, alpha6, alpha8, alpha10. The radial grid is cached so only the five scalar coefficients are optimized.

Attributes:

Name Type Description
alpha2 Tensor

Coefficient of \(r^2\). Scalar tensor, shape [1].

alpha4 Tensor

Coefficient of \(r^4\). Scalar tensor, shape [1].

alpha6 Tensor

Coefficient of \(r^6\). Scalar tensor, shape [1].

alpha8 Tensor

Coefficient of \(r^8\). Scalar tensor, shape [1].

alpha10 Tensor

Coefficient of \(r^{10}\). Scalar tensor, shape [1].

x Tensor

Pixel x-coordinates. [H, W]. [mm]

y Tensor

Pixel y-coordinates. [H, W]. [mm]

r2 Tensor

Cached squared radius \(x^2 + y^2\). [H, W]. [mm^2]

Initialize a Binary2 DOE with small random polynomial coefficients.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
res tuple or int

Resolution as (H, W); an int is expanded to (res, res). [pixel]. Defaults to (2000, 2000).

(2000, 2000)
mat str

DOE material name. Defaults to "fused_silica".

'fused_silica'
wvln0 float

Design wavelength. [um]. Defaults to 0.55.

0.55
fab_ps float

Fabrication pixel size. [mm]. Defaults to 0.001.

0.001
fab_step int

Number of fabrication quantization levels. Defaults to 16.

16
is_square bool

Whether the aperture is square. Defaults to True.

True
device str

Device to store tensors on. Defaults to "cpu".

'cpu'
Source code in deeplens-src/deeplens/diffractive_surface/binary2.py
def __init__(
    self,
    d_next,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    is_square=True,
    device="cpu",
):
    """Initialize a Binary2 DOE with small random polynomial coefficients.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        res (tuple or int, optional): Resolution as (H, W); an int is
            expanded to (res, res). [pixel]. Defaults to (2000, 2000).
        mat (str, optional): DOE material name. Defaults to "fused_silica".
        wvln0 (float, optional): Design wavelength. [um]. Defaults to 0.55.
        fab_ps (float, optional): Fabrication pixel size. [mm]. Defaults to 0.001.
        fab_step (int, optional): Number of fabrication quantization levels. Defaults to 16.
        is_square (bool, optional): Whether the aperture is square. Defaults to True.
        device (str, optional): Device to store tensors on. Defaults to "cpu".
    """
    super().__init__(
        d_next=d_next,
        res=res,
        mat=mat,
        wvln0=wvln0,
        fab_ps=fab_ps,
        fab_step=fab_step,
        is_square=is_square,
        device=device,
    )

    # Initialize with random small values
    self.alpha2 = (torch.rand(1) - 0.5) * 0.02
    self.alpha4 = (torch.rand(1) - 0.5) * 0.002
    self.alpha6 = (torch.rand(1) - 0.5) * 0.0002
    self.alpha8 = (torch.rand(1) - 0.5) * 0.00002
    self.alpha10 = (torch.rand(1) - 0.5) * 0.000002

    self.x, self.y = torch.meshgrid(
        torch.linspace(-self.w / 2, self.w / 2, self.res[1]),
        torch.linspace(self.h / 2, -self.h / 2, self.res[0]),
        indexing="xy",
    )

    # Cache static r² grid (x, y never change after init)
    self.r2 = self.x**2 + self.y**2

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a Binary2 DOE from a serialized surface dict.

Parameters:

Name Type Description Default
doe_dict dict

Surface dict. Requires keys "d_next" and "res"; optional keys "mat", "wvln0", "fab_ps", "fab_step", "is_square" fall back to the constructor defaults.

required

Returns:

Name Type Description
doe Binary2

The constructed Binary2 surface.

Source code in deeplens-src/deeplens/diffractive_surface/binary2.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a Binary2 DOE from a serialized surface dict.

    Args:
        doe_dict (dict): Surface dict. Requires keys "d_next" and "res"; optional
            keys "mat", "wvln0", "fab_ps", "fab_step", "is_square" fall back
            to the constructor defaults.

    Returns:
        doe (Binary2): The constructed Binary2 surface.
    """
    return cls(
        d_next=doe_dict["d_next"],
        res=doe_dict["res"],
        mat=doe_dict.get("mat", "fused_silica"),
        wvln0=doe_dict.get("wvln0", 0.55),
        fab_ps=doe_dict.get("fab_ps", 0.001),
        fab_step=doe_dict.get("fab_step", 16),
        is_square=doe_dict.get("is_square", True),
    )

phase_func

phase_func()

Compute the raw (unwrapped) phase at the design wavelength.

Evaluates \(\phi(r) = \pi\,(\alpha_2 r^2 + \alpha_4 r^4 + \alpha_6 r^6 + \alpha_8 r^8 + \alpha_{10} r^{10})\) via Horner's method on the cached \(r^2\) grid.

Returns:

Name Type Description
phase Tensor

Raw phase map. [H, W]. [rad]

Source code in deeplens-src/deeplens/diffractive_surface/binary2.py
def phase_func(self):
    """Compute the raw (unwrapped) phase at the design wavelength.

    Evaluates $\\phi(r) = \\pi\\,(\\alpha_2 r^2 + \\alpha_4 r^4 + \\alpha_6 r^6
    + \\alpha_8 r^8 + \\alpha_{10} r^{10})$ via Horner's method on the cached
    $r^2$ grid.

    Returns:
        phase (torch.Tensor): Raw phase map. [H, W]. [rad]
    """
    # Horner's method: r2*(a2 + r2*(a4 + r2*(a6 + r2*(a8 + r2*a10))))
    r2 = self.r2
    phase = (
        torch.pi
        * r2
        * (
            self.alpha2
            + r2
            * (
                self.alpha4
                + r2 * (self.alpha6 + r2 * (self.alpha8 + r2 * self.alpha10))
            )
        )
    )
    return phase

get_optimizer_params

get_optimizer_params(lr=0.001)

Enable gradients and build per-coefficient optimizer parameter groups.

Higher-order coefficients use progressively larger learning rates (lr, 10x, 100x, 1000x, 10000x for alpha2 through alpha10) to compensate for their smaller magnitude.

Parameters:

Name Type Description Default
lr float

Base learning rate for alpha2. Defaults to 0.001.

0.001

Returns:

Name Type Description
optimizer_params list

List of parameter-group dicts, one per coefficient, each with keys "params" and "lr".

Source code in deeplens-src/deeplens/diffractive_surface/binary2.py
def get_optimizer_params(self, lr=0.001):
    """Enable gradients and build per-coefficient optimizer parameter groups.

    Higher-order coefficients use progressively larger learning rates
    (`lr`, 10x, 100x, 1000x, 10000x for `alpha2` through `alpha10`) to
    compensate for their smaller magnitude.

    Args:
        lr (float): Base learning rate for `alpha2`. Defaults to 0.001.

    Returns:
        optimizer_params (list): List of parameter-group dicts, one per
            coefficient, each with keys "params" and "lr".
    """
    self.alpha2.requires_grad = True
    self.alpha4.requires_grad = True
    self.alpha6.requires_grad = True
    self.alpha8.requires_grad = True
    self.alpha10.requires_grad = True

    optimizer_params = [
        {"params": [self.alpha2], "lr": lr},
        {"params": [self.alpha4], "lr": lr * 10},
        {"params": [self.alpha6], "lr": lr * 100},
        {"params": [self.alpha8], "lr": lr * 1000},
        {"params": [self.alpha10], "lr": lr * 10000},
    ]

    return optimizer_params

surf_dict

surf_dict()

Serialize the surface to a dict, including the polynomial coefficients.

Returns:

Name Type Description
surf_dict dict

Base surface dict extended with the five rounded coefficients "alpha2", "alpha4", "alpha6", "alpha8", "alpha10".

Source code in deeplens-src/deeplens/diffractive_surface/binary2.py
def surf_dict(self):
    """Serialize the surface to a dict, including the polynomial coefficients.

    Returns:
        surf_dict (dict): Base surface dict extended with the five rounded
            coefficients "alpha2", "alpha4", "alpha6", "alpha8", "alpha10".
    """
    surf_dict = super().surf_dict()
    surf_dict["alpha2"] = round(self.alpha2.item(), 6)
    surf_dict["alpha4"] = round(self.alpha4.item(), 6)
    surf_dict["alpha6"] = round(self.alpha6.item(), 6)
    surf_dict["alpha8"] = round(self.alpha8.item(), 6)
    surf_dict["alpha10"] = round(self.alpha10.item(), 6)
    return surf_dict

Free-form, per-pixel phase map.

deeplens.diffractive_surface.Pixel2D

Pixel2D(
    d_next,
    phase_map_path=None,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    device="cpu",
)

Bases: DiffractiveSurface

Pixel2D DOE parameterization with a direct, per-pixel phase map.

Each pixel of the phase map is an independent optimizable parameter, giving the most general (and highest-dimensional) DOE parameterization. The phase map is stored at the design wavelength wvln0.

Attributes:

Name Type Description
phase_map Tensor

Per-pixel phase at the design wavelength. [H, W]. [rad]

Initialize a Pixel2D DOE where each pixel is an independent parameter.

If phase_map_path is None the phase map is initialized to small random values (torch.randn * 1e-3); otherwise it is loaded from the given path.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
phase_map_path str or None

Path to a saved phase-map tensor to load. If None, the phase map is randomly initialized. Defaults to None.

None
res tuple or int

Resolution of the DOE as (H, W); an int is expanded to (res, res). [pixel]. Defaults to (2000, 2000).

(2000, 2000)
mat str

Material of the DOE. Defaults to "fused_silica".

'fused_silica'
wvln0 float

Design wavelength. [um]. Defaults to 0.55.

0.55
fab_ps float

Fabrication pixel size. [mm]. Defaults to 0.001.

0.001
fab_step int

Number of fabrication quantization levels. Defaults to 16.

16
device str

Device to run the DOE. Defaults to "cpu".

'cpu'

Raises:

Type Description
ValueError

If phase_map_path is neither None nor a string.

Source code in deeplens-src/deeplens/diffractive_surface/pixel2d.py
def __init__(
    self,
    d_next,
    phase_map_path=None,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    device="cpu",
):
    """Initialize a Pixel2D DOE where each pixel is an independent parameter.

    If `phase_map_path` is None the phase map is initialized to small random
    values (`torch.randn * 1e-3`); otherwise it is loaded from the given path.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        phase_map_path (str or None, optional): Path to a saved phase-map
            tensor to load. If None, the phase map is randomly initialized.
            Defaults to None.
        res (tuple or int, optional): Resolution of the DOE as (H, W); an int
            is expanded to (res, res). [pixel]. Defaults to (2000, 2000).
        mat (str, optional): Material of the DOE. Defaults to "fused_silica".
        wvln0 (float, optional): Design wavelength. [um]. Defaults to 0.55.
        fab_ps (float, optional): Fabrication pixel size. [mm]. Defaults to 0.001.
        fab_step (int, optional): Number of fabrication quantization levels.
            Defaults to 16.
        device (str, optional): Device to run the DOE. Defaults to "cpu".

    Raises:
        ValueError: If `phase_map_path` is neither None nor a string.
    """
    super().__init__(
        d_next=d_next,
        res=res,
        mat=mat,
        fab_ps=fab_ps,
        fab_step=fab_step,
        wvln0=wvln0,
        device=device,
    )

    # Initialize phase map with random values
    if phase_map_path is None:
        self.phase_map = torch.randn(self.res, device=self.device) * 1e-3
    elif isinstance(phase_map_path, str):
        self.phase_map = torch.load(
            phase_map_path, map_location=device, weights_only=True
        )
    else:
        raise ValueError(f"Invalid phase_map_path: {phase_map_path}")

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a Pixel2D DOE from a dict.

Parameters:

Name Type Description Default
doe_dict dict

Surface dict with keys "d_next" and "res" required and optional keys "mat", "fab_ps", "fab_step", "phase_map_path", "wvln0".

required

Returns:

Name Type Description
doe Pixel2D

The constructed Pixel2D DOE.

Source code in deeplens-src/deeplens/diffractive_surface/pixel2d.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a Pixel2D DOE from a dict.

    Args:
        doe_dict (dict): Surface dict with keys "d_next" and "res" required and
            optional keys "mat", "fab_ps", "fab_step", "phase_map_path", "wvln0".

    Returns:
        doe (Pixel2D): The constructed Pixel2D DOE.
    """
    return cls(
        d_next=doe_dict["d_next"],
        res=doe_dict["res"],
        mat=doe_dict.get("mat", "fused_silica"),
        fab_ps=doe_dict.get("fab_ps", 0.001),
        fab_step=doe_dict.get("fab_step", 16),
        phase_map_path=doe_dict.get("phase_map_path", None),
        wvln0=doe_dict.get("wvln0", 0.55),
    )

phase_func

phase_func()

Return the raw per-pixel phase map at the design wavelength.

Returns:

Name Type Description
phase_map Tensor

Per-pixel phase at the design wavelength. [H, W]. [rad]

Source code in deeplens-src/deeplens/diffractive_surface/pixel2d.py
def phase_func(self):
    """Return the raw per-pixel phase map at the design wavelength.

    Returns:
        phase_map (torch.Tensor): Per-pixel phase at the design wavelength.
            [H, W]. [rad]
    """
    return self.phase_map

get_optimizer_params

get_optimizer_params(lr=0.01)

Get optimizer parameter groups for the phase map.

Enables gradients on the phase map and returns it as a single Adam-style parameter group with the given learning rate.

Parameters:

Name Type Description Default
lr float

Learning rate for the phase map. Defaults to 0.01.

0.01

Returns:

Name Type Description
optimizer_params list

List with one parameter-group dict {"params": [phase_map], "lr": lr}.

Source code in deeplens-src/deeplens/diffractive_surface/pixel2d.py
def get_optimizer_params(self, lr=0.01):
    """Get optimizer parameter groups for the phase map.

    Enables gradients on the phase map and returns it as a single Adam-style
    parameter group with the given learning rate.

    Args:
        lr (float, optional): Learning rate for the phase map. Defaults to 0.01.

    Returns:
        optimizer_params (list): List with one parameter-group dict
            {"params": [phase_map], "lr": lr}.
    """
    self.phase_map.requires_grad = True
    optimizer_params = [{"params": [self.phase_map], "lr": lr}]
    return optimizer_params

surf_dict

surf_dict(phase_map_path)

Return a serializable surface dict and save the phase map to disk.

Extends the base surface dict with the phase-map path, and writes the detached CPU phase-map tensor to phase_map_path.

Parameters:

Name Type Description Default
phase_map_path str

Path to which the phase-map tensor is saved and which is recorded in the returned dict.

required

Returns:

Name Type Description
surf_dict dict

Surface dict including the "phase_map_path" entry.

Source code in deeplens-src/deeplens/diffractive_surface/pixel2d.py
def surf_dict(self, phase_map_path):
    """Return a serializable surface dict and save the phase map to disk.

    Extends the base surface dict with the phase-map path, and writes the
    detached CPU phase-map tensor to `phase_map_path`.

    Args:
        phase_map_path (str): Path to which the phase-map tensor is saved and
            which is recorded in the returned dict.

    Returns:
        surf_dict (dict): Surface dict including the "phase_map_path" entry.
    """
    surf_dict = super().surf_dict()
    surf_dict["phase_map_path"] = phase_map_path
    torch.save(self.phase_map.clone().detach().cpu(), phase_map_path)
    return surf_dict

Fresnel-lens (quadratic) phase profile.

deeplens.diffractive_surface.Fresnel

Fresnel(
    d_next,
    f0=None,
    wvln0=0.55,
    res=(2000, 2000),
    mat="fused_silica",
    fab_ps=0.001,
    fab_step=16,
    device="cpu",
)

Bases: DiffractiveSurface

Phase-Fresnel diffractive lens surface.

A diffractive Fresnel lens with an ideal quadratic (thin-lens) phase profile. It exhibits inverse dispersion compared to a refractive lens, and its only free parameter is the design-wavelength focal length f0.

Attributes:

Name Type Description
f0 Tensor

Design-wavelength focal length, scalar. [mm]

r2 Tensor

Cached squared radial coordinate grid \(x^2 + y^2\), shape [H, W]. [mm^2]

Initialize a phase-Fresnel diffractive lens.

The lens applies an ideal thin-lens quadratic phase set by f0. It shows inverse dispersion compared to a refractive lens.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
f0 float or None

Design-wavelength focal length. [mm] If None, initialized to a random near-infinite value. Defaults to None.

None
wvln0 float

Design wavelength. [um] Defaults to 0.55.

0.55
res tuple or int

Resolution of the DOE, [w, h]. [pixel] Defaults to (2000, 2000).

(2000, 2000)
mat str

Material of the DOE. Defaults to "fused_silica".

'fused_silica'
fab_ps float

Fabrication pixel size. [mm] Defaults to 0.001.

0.001
fab_step int

Number of fabrication quantization steps. Defaults to 16.

16
device str

Device to run the DOE. Defaults to "cpu".

'cpu'
Source code in deeplens-src/deeplens/diffractive_surface/fresnel.py
def __init__(
    self,
    d_next,
    f0=None,
    wvln0=0.55,
    res=(2000, 2000),
    mat="fused_silica",
    fab_ps=0.001,
    fab_step=16,
    device="cpu",
):
    """Initialize a phase-Fresnel diffractive lens.

    The lens applies an ideal thin-lens quadratic phase set by `f0`. It shows
    inverse dispersion compared to a refractive lens.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        f0 (float or None, optional): Design-wavelength focal length. [mm]
            If None, initialized to a random near-infinite value. Defaults to None.
        wvln0 (float, optional): Design wavelength. [um] Defaults to 0.55.
        res (tuple or int, optional): Resolution of the DOE, [w, h]. [pixel]
            Defaults to (2000, 2000).
        mat (str, optional): Material of the DOE. Defaults to "fused_silica".
        fab_ps (float, optional): Fabrication pixel size. [mm] Defaults to 0.001.
        fab_step (int, optional): Number of fabrication quantization steps.
            Defaults to 16.
        device (str, optional): Device to run the DOE. Defaults to "cpu".
    """
    super().__init__(
        d_next=d_next,
        res=res,
        wvln0=wvln0,
        mat=mat,
        fab_ps=fab_ps,
        fab_step=fab_step,
        device=device,
    )

    # Initial focal length
    if f0 is None:
        self.f0 = torch.randn(1) * 1e6
    else:
        self.f0 = torch.tensor(f0)

    # Cache static r² grid (x, y never change after init)
    self.r2 = self.x**2 + self.y**2

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a Fresnel DOE from a dictionary of surface parameters.

Parameters:

Name Type Description Default
doe_dict dict

Surface parameters. Requires "d_next" and "res"; optionally "f0", "wvln0", "mat", "fab_ps", "fab_step".

required

Returns:

Name Type Description
doe Fresnel

The constructed Fresnel DOE.

Source code in deeplens-src/deeplens/diffractive_surface/fresnel.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a Fresnel DOE from a dictionary of surface parameters.

    Args:
        doe_dict (dict): Surface parameters. Requires "d_next" and "res"; optionally
            "f0", "wvln0", "mat", "fab_ps", "fab_step".

    Returns:
        doe (Fresnel): The constructed Fresnel DOE.
    """
    return cls(
        d_next=doe_dict["d_next"],
        res=doe_dict["res"],
        fab_ps=doe_dict.get("fab_ps", 0.001),
        fab_step=doe_dict.get("fab_step", 16),
        f0=doe_dict.get("f0", None),
        wvln0=doe_dict.get("wvln0", 0.55),
        mat=doe_dict.get("mat", "fused_silica"),
    )

phase_func

phase_func()

Compute the raw (unwrapped) quadratic phase at the design wavelength.

Applies the ideal thin-lens phase

\[\phi(x, y) = -\frac{\pi (x^2 + y^2)}{f_0 \lambda_0}\]

where \(\lambda_0\) is the design wavelength converted to mm. Emits a one-time warning if the phase is undersampled on the current grid.

Returns:

Name Type Description
phase Tensor

Raw unwrapped phase, shape [H, W]. [rad]

Source code in deeplens-src/deeplens/diffractive_surface/fresnel.py
def phase_func(self):
    """Compute the raw (unwrapped) quadratic phase at the design wavelength.

    Applies the ideal thin-lens phase

    $$\\phi(x, y) = -\\frac{\\pi (x^2 + y^2)}{f_0 \\lambda_0}$$

    where $\\lambda_0$ is the design wavelength converted to mm. Emits a
    one-time warning if the phase is undersampled on the current grid.

    Returns:
        phase (torch.Tensor): Raw unwrapped phase, shape [H, W]. [rad]
    """
    wvln0_mm = self.wvln0 * 1e-3
    phase = -2 * torch.pi * self.r2 / (2 * self.f0 * wvln0_mm)
    self._warn_if_undersampled(phase, self.f0, self.wvln0)
    return phase

get_optimizer_params

get_optimizer_params(lr=0.001)

Build optimizer parameter groups for the focal length f0.

Enables gradients on f0 and returns it as a single parameter group.

Parameters:

Name Type Description Default
lr float

Learning rate for f0. Defaults to 0.001.

0.001

Returns:

Name Type Description
optimizer_params list

List with one parameter group dict for f0.

Source code in deeplens-src/deeplens/diffractive_surface/fresnel.py
def get_optimizer_params(self, lr=0.001):
    """Build optimizer parameter groups for the focal length `f0`.

    Enables gradients on `f0` and returns it as a single parameter group.

    Args:
        lr (float, optional): Learning rate for `f0`. Defaults to 0.001.

    Returns:
        optimizer_params (list): List with one parameter group dict for `f0`.
    """
    self.f0.requires_grad = True
    optimizer_params = [{"params": [self.f0], "lr": lr}]
    return optimizer_params

surf_dict

surf_dict()

Serialize the surface to a dictionary, including f0 and wvln0.

Returns:

Name Type Description
surf_dict dict

Base surface parameters plus "f0" [mm], with "wvln0" [um] overwritten by the unrounded value.

Source code in deeplens-src/deeplens/diffractive_surface/fresnel.py
def surf_dict(self):
    """Serialize the surface to a dictionary, including `f0` and `wvln0`.

    Returns:
        surf_dict (dict): Base surface parameters plus "f0" [mm], with
            "wvln0" [um] overwritten by the unrounded value.
    """
    surf_dict = super().surf_dict()
    surf_dict["f0"] = self.f0.item()
    surf_dict["wvln0"] = self.wvln0
    return surf_dict

Phase parameterized by Zernike polynomials.

deeplens.diffractive_surface.Zernike

Zernike(
    d_next,
    z_coeff=None,
    zernike_order=37,
    res=(2000, 2000),
    mat="fused_silica",
    fab_ps=0.001,
    fab_step=16,
    wvln0=0.55,
    device="cpu",
)

Bases: DiffractiveSurface

Diffractive optical element parameterized by Zernike polynomials.

The DOE surface phase is represented as a weighted sum of the first 37 Zernike polynomials (OSA/ANSI ordering) over the unit disk. The learnable coefficients z_coeff are the only optimized parameters.

Attributes:

Name Type Description
zernike_order int

Number of Zernike terms (fixed at 37).

z_coeff Tensor

Zernike coefficients, shape (zernike_order,).

Initialize a Zernike-parameterized DOE.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
z_coeff Tensor or None

Zernike coefficients of shape (zernike_order,). If None, initialized to random values scaled by 1e-3. Defaults to None.

None
zernike_order int

Number of Zernike coefficients. Only 37 is currently supported. Defaults to 37.

37
res tuple

DOE resolution as (H, W) in pixels. Defaults to (2000, 2000).

(2000, 2000)
mat str

DOE substrate material. Defaults to "fused_silica".

'fused_silica'
fab_ps float

Fabrication pixel size. [mm] Defaults to 0.001.

0.001
fab_step int

Number of fabrication quantization levels. Defaults to 16.

16
wvln0 float

Design wavelength. [um] Defaults to 0.55.

0.55
device str

Computation device. Defaults to "cpu".

'cpu'

Raises:

Type Description
AssertionError

If zernike_order is not 37.

Source code in deeplens-src/deeplens/diffractive_surface/zernike.py
def __init__(
    self,
    d_next,
    z_coeff=None,
    zernike_order=37,
    res=(2000, 2000),
    mat="fused_silica",
    fab_ps=0.001,
    fab_step=16,
    wvln0=0.55,
    device="cpu",
):
    """Initialize a Zernike-parameterized DOE.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        z_coeff (torch.Tensor or None, optional): Zernike coefficients of
            shape (zernike_order,). If None, initialized to random values
            scaled by 1e-3. Defaults to None.
        zernike_order (int, optional): Number of Zernike coefficients. Only
            37 is currently supported. Defaults to 37.
        res (tuple, optional): DOE resolution as (H, W) in pixels. Defaults
            to (2000, 2000).
        mat (str, optional): DOE substrate material. Defaults to "fused_silica".
        fab_ps (float, optional): Fabrication pixel size. [mm] Defaults to 0.001.
        fab_step (int, optional): Number of fabrication quantization levels.
            Defaults to 16.
        wvln0 (float, optional): Design wavelength. [um] Defaults to 0.55.
        device (str, optional): Computation device. Defaults to "cpu".

    Raises:
        AssertionError: If zernike_order is not 37.
    """
    super().__init__(
        d_next=d_next,
        res=res,
        mat=mat,
        fab_ps=fab_ps,
        fab_step=fab_step,
        wvln0=wvln0,
        device=device,
    )

    # Initialize Zernike coefficients with random values
    assert zernike_order == 37, "Currently, Zernike DOE only supports 37 orders"
    self.zernike_order = zernike_order
    if z_coeff is None:
        self.z_coeff = torch.randn(zernike_order, device=self.device) * 1e-3
    else:
        self.z_coeff = z_coeff

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a Zernike DOE from a serialized surface dict.

Parameters:

Name Type Description Default
doe_dict dict

Surface parameters. Requires "d_next" and "res"; optional keys "mat", "fab_ps", "fab_step", "z_coeff", "zernike_order", "wvln0" fall back to their defaults when absent.

required

Returns:

Name Type Description
zernike Zernike

The constructed Zernike DOE.

Source code in deeplens-src/deeplens/diffractive_surface/zernike.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a Zernike DOE from a serialized surface dict.

    Args:
        doe_dict (dict): Surface parameters. Requires "d_next" and "res"; optional
            keys "mat", "fab_ps", "fab_step", "z_coeff", "zernike_order",
            "wvln0" fall back to their defaults when absent.

    Returns:
        zernike (Zernike): The constructed Zernike DOE.
    """
    return cls(
        d_next=doe_dict["d_next"],
        res=doe_dict["res"],
        mat=doe_dict.get("mat", "fused_silica"),
        fab_ps=doe_dict.get("fab_ps", 0.001),
        fab_step=doe_dict.get("fab_step", 16),
        z_coeff=doe_dict.get("z_coeff", None),
        zernike_order=doe_dict.get("zernike_order", 37),
        wvln0=doe_dict.get("wvln0", 0.55),
    )

phase_func

phase_func()

Compute the DOE phase map at the design wavelength.

Returns:

Name Type Description
phase Tensor

Phase map of shape (res[0], res[0]) in radians, evaluated from the Zernike coefficients over the unit disk.

Source code in deeplens-src/deeplens/diffractive_surface/zernike.py
def phase_func(self):
    """Compute the DOE phase map at the design wavelength.

    Returns:
        phase (torch.Tensor): Phase map of shape (res[0], res[0]) in radians,
            evaluated from the Zernike coefficients over the unit disk.
    """
    return calculate_zernike_phase(self.z_coeff, grid=self.res[0])

get_optimizer_params

get_optimizer_params(lr=0.01)

Build optimizer parameter groups for the Zernike coefficients.

Sets z_coeff to require gradients as a side effect.

Parameters:

Name Type Description Default
lr float

Learning rate for the coefficients. Defaults to 0.01.

0.01

Returns:

Name Type Description
optimizer_params list

A single parameter group dict with keys "params" (the z_coeff tensor) and "lr".

Source code in deeplens-src/deeplens/diffractive_surface/zernike.py
def get_optimizer_params(self, lr=0.01):
    """Build optimizer parameter groups for the Zernike coefficients.

    Sets `z_coeff` to require gradients as a side effect.

    Args:
        lr (float, optional): Learning rate for the coefficients. Defaults to 0.01.

    Returns:
        optimizer_params (list): A single parameter group dict with keys
            "params" (the `z_coeff` tensor) and "lr".
    """
    self.z_coeff.requires_grad = True
    optimizer_params = [{"params": [self.z_coeff], "lr": lr}]
    return optimizer_params

surf_dict

surf_dict()

Serialize the DOE surface to a dict.

Extends the base surface dict with the Zernike coefficients (moved to CPU and detached) and the Zernike order.

Returns:

Name Type Description
surf_dict dict

Surface parameters including "z_coeff" and "zernike_order".

Source code in deeplens-src/deeplens/diffractive_surface/zernike.py
def surf_dict(self):
    """Serialize the DOE surface to a dict.

    Extends the base surface dict with the Zernike coefficients (moved to
    CPU and detached) and the Zernike order.

    Returns:
        surf_dict (dict): Surface parameters including "z_coeff" and
            "zernike_order".
    """
    surf_dict = super().surf_dict()
    surf_dict["z_coeff"] = self.z_coeff.clone().detach().cpu()
    surf_dict["zernike_order"] = self.zernike_order
    return surf_dict

Linear / blazed grating phase.

deeplens.diffractive_surface.Grating

Grating(
    d_next,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    theta=0.0,
    alpha=0.0,
    device="cpu",
)

Bases: DiffractiveSurface

Linear grating diffractive optical element.

A grating introduces a linear phase gradient across the surface, which diffracts light into multiple diffraction orders. The phase profile is

\[\phi(x, y) = \alpha \,\frac{x \sin\theta + y \cos\theta}{\text{norm\_radii}}\]

where \(\theta\) is the angle from the y-axis to the grating vector, \(\alpha\) is the grating slope (phase-gradient strength), and norm_radii normalizes the coordinates.

Attributes:

Name Type Description
theta Tensor

Angle from the y-axis to the grating vector. [rad]

alpha Tensor

Grating slope (phase-gradient strength). [rad]

norm_radii float

Coordinate normalization radius (half the DOE width). [mm]

Initialize a grating DOE.

Parameters:

Name Type Description Default
d_next float

Axial thickness to the next plane. [mm]

required
res tuple or int

Resolution of the DOE as (H, W); an int is expanded to (res, res). [pixel]. Defaults to (2000, 2000).

(2000, 2000)
mat str

Material name of the DOE. Defaults to "fused_silica".

'fused_silica'
wvln0 float

Design wavelength. [um]. Defaults to 0.55.

0.55
fab_ps float

Fabrication pixel size. [mm]. Defaults to 0.001.

0.001
fab_step int

Number of fabrication (quantization) levels. Defaults to 16.

16
theta float

Angle from the y-axis to the grating vector. [rad]. Defaults to 0.0.

0.0
alpha float

Grating slope (phase-gradient strength). [rad]. Defaults to 0.0.

0.0
device str

Device to place the DOE tensors on. Defaults to "cpu".

'cpu'
Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def __init__(
    self,
    d_next,
    res=(2000, 2000),
    mat="fused_silica",
    wvln0=0.55,
    fab_ps=0.001,
    fab_step=16,
    theta=0.0,
    alpha=0.0,
    device="cpu",
):
    """Initialize a grating DOE.

    Args:
        d_next (float): Axial thickness to the next plane. [mm]
        res (tuple or int, optional): Resolution of the DOE as (H, W); an
            int is expanded to (res, res). [pixel]. Defaults to (2000, 2000).
        mat (str, optional): Material name of the DOE. Defaults to "fused_silica".
        wvln0 (float, optional): Design wavelength. [um]. Defaults to 0.55.
        fab_ps (float, optional): Fabrication pixel size. [mm]. Defaults to 0.001.
        fab_step (int, optional): Number of fabrication (quantization)
            levels. Defaults to 16.
        theta (float, optional): Angle from the y-axis to the grating
            vector. [rad]. Defaults to 0.0.
        alpha (float, optional): Grating slope (phase-gradient strength).
            [rad]. Defaults to 0.0.
        device (str, optional): Device to place the DOE tensors on. Defaults to "cpu".
    """
    super().__init__(
        d_next=d_next,
        res=res,
        mat=mat,
        wvln0=wvln0,
        fab_ps=fab_ps,
        fab_step=fab_step,
        device=device,
    )

    # Grating parameters
    self.theta = torch.tensor(theta)  # angle from y-axis to grating vector
    self.alpha = torch.tensor(alpha)  # slope of the grating

    # Normalization radius (use half of the width)
    self.norm_radii = self.w / 2

    self.to(device)

init_from_dict classmethod

init_from_dict(doe_dict)

Initialize a grating DOE from a parameter dict.

Parameters:

Name Type Description Default
doe_dict dict

Dictionary of DOE parameters. Requires keys "d_next" and "res"; "mat", "wvln0", "fab_ps", "fab_step", "theta", and "alpha" are optional and fall back to defaults.

required

Returns:

Name Type Description
grating Grating

The constructed grating DOE instance.

Source code in deeplens-src/deeplens/diffractive_surface/grating.py
@classmethod
def init_from_dict(cls, doe_dict):
    """Initialize a grating DOE from a parameter dict.

    Args:
        doe_dict (dict): Dictionary of DOE parameters. Requires keys "d_next" and
            "res"; "mat", "wvln0", "fab_ps", "fab_step", "theta", and
            "alpha" are optional and fall back to defaults.

    Returns:
        grating (Grating): The constructed grating DOE instance.
    """
    return cls(
        d_next=doe_dict["d_next"],
        res=doe_dict["res"],
        mat=doe_dict.get("mat", "fused_silica"),
        wvln0=doe_dict.get("wvln0", 0.55),
        fab_ps=doe_dict.get("fab_ps", 0.001),
        fab_step=doe_dict.get("fab_step", 16),
        theta=doe_dict.get("theta", 0.0),
        alpha=doe_dict.get("alpha", 0.0),
    )

phase_func

phase_func()

Compute the raw grating phase profile at the design wavelength.

The phase is a linear function of position:

\[\phi(x, y) = \alpha \,\frac{x \sin\theta + y \cos\theta}{\text{norm\_radii}}\]

Returns:

Name Type Description
phase Tensor

Raw, unwrapped phase profile at the design wavelength. [H, W]. [rad]

Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def phase_func(self):
    """Compute the raw grating phase profile at the design wavelength.

    The phase is a linear function of position:

    $$\\phi(x, y) = \\alpha \\,\\frac{x \\sin\\theta + y \\cos\\theta}{\\text{norm\\_radii}}$$

    Returns:
        phase (torch.Tensor): Raw, unwrapped phase profile at the design
            wavelength. [H, W]. [rad]
    """
    # Normalize coordinates
    x_norm = self.x / self.norm_radii
    y_norm = self.y / self.norm_radii

    # Calculate linear phase gradient
    phase = self.alpha * (
        x_norm * torch.sin(self.theta) + y_norm * torch.cos(self.theta)
    )

    return phase

get_optimizer_params

get_optimizer_params(lr=0.001)

Build optimizer parameter groups for the grating parameters.

Enables gradients on theta and alpha. The alpha group uses a learning rate scaled by 10x relative to lr.

Parameters:

Name Type Description Default
lr float

Base learning rate for the grating parameters. Defaults to 0.001.

0.001

Returns:

Name Type Description
optimizer_params list

List of parameter-group dicts for the optimizer.

Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def get_optimizer_params(self, lr=0.001):
    """Build optimizer parameter groups for the grating parameters.

    Enables gradients on `theta` and `alpha`. The `alpha` group uses a
    learning rate scaled by 10x relative to `lr`.

    Args:
        lr (float, optional): Base learning rate for the grating
            parameters. Defaults to 0.001.

    Returns:
        optimizer_params (list): List of parameter-group dicts for the
            optimizer.
    """
    self.theta.requires_grad = True
    self.alpha.requires_grad = True

    optimizer_params = [
        {"params": [self.theta], "lr": lr},
        {"params": [self.alpha], "lr": lr * 10},
    ]

    return optimizer_params

surf_dict

surf_dict()

Return a serializable dict of the grating surface parameters.

Extends the base surface dict with the grating-specific theta, alpha, and norm_radii entries.

Returns:

Name Type Description
surf_dict dict

Dictionary of surface parameters.

Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def surf_dict(self):
    """Return a serializable dict of the grating surface parameters.

    Extends the base surface dict with the grating-specific `theta`,
    `alpha`, and `norm_radii` entries.

    Returns:
        surf_dict (dict): Dictionary of surface parameters.
    """
    surf_dict = super().surf_dict()
    surf_dict["theta"] = round(self.theta.item(), 6)
    surf_dict["alpha"] = round(self.alpha.item(), 6)
    surf_dict["norm_radii"] = round(self.norm_radii, 6)
    return surf_dict

save_ckpt

save_ckpt(save_path='./grating_doe.pth')

Save the grating DOE parameters to a checkpoint file.

Parameters:

Name Type Description Default
save_path str

Path to write the checkpoint to. Defaults to "./grating_doe.pth".

'./grating_doe.pth'
Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def save_ckpt(self, save_path="./grating_doe.pth"):
    """Save the grating DOE parameters to a checkpoint file.

    Args:
        save_path (str, optional): Path to write the checkpoint to. Defaults
            to "./grating_doe.pth".
    """
    torch.save(
        {
            "param_model": "grating",
            "theta": self.theta.clone().detach().cpu(),
            "alpha": self.alpha.clone().detach().cpu(),
        },
        save_path,
    )

load_ckpt

load_ckpt(load_path='./grating_doe.pth')

Load the grating DOE parameters from a checkpoint file.

Restores theta and alpha onto the current device.

Parameters:

Name Type Description Default
load_path str

Path to read the checkpoint from. Defaults to "./grating_doe.pth".

'./grating_doe.pth'
Source code in deeplens-src/deeplens/diffractive_surface/grating.py
def load_ckpt(self, load_path="./grating_doe.pth"):
    """Load the grating DOE parameters from a checkpoint file.

    Restores `theta` and `alpha` onto the current device.

    Args:
        load_path (str, optional): Path to read the checkpoint from.
            Defaults to "./grating_doe.pth".
    """
    ckpt = torch.load(load_path)
    self.theta = ckpt["theta"].to(self.device)
    self.alpha = ckpt["alpha"].to(self.device)