Skip to content

Surrogate Networks

Neural networks that learn to predict PSFs from lens parameters, replacing expensive ray tracing during training. These power PSFNetLens.

Fully-connected network that predicts PSF values from input parameters.

deeplens.surrogate.MLP

MLP(in_features, out_features, hidden_features=64, hidden_layers=3)

Bases: Module

Fully-connected network for low-frequency PSF prediction.

Predicts PSFs as flattened vectors using stacked linear layers with ReLU activations and a Sigmoid output. The output is L1-normalized so it sums to 1 (valid as a PSF energy distribution).

Parameters:

Name Type Description Default
in_features int

Number of input features (e.g., field angle + wavelength).

required
out_features int

Number of output features (flattened PSF size).

required
hidden_features int

Width of hidden layers. Defaults to 64.

64
hidden_layers int

Number of hidden layers. Defaults to 3.

3
Source code in deeplens-src/deeplens/surrogate/mlp.py
def __init__(self, in_features, out_features, hidden_features=64, hidden_layers=3):
    super(MLP, self).__init__()

    layers = [
        nn.Linear(in_features, hidden_features // 4, bias=True),
        nn.ReLU(inplace=True),
        nn.Linear(hidden_features // 4, hidden_features, bias=True),
        nn.ReLU(inplace=True),
    ]

    for _ in range(hidden_layers):
        layers.extend(
            [
                nn.Linear(hidden_features, hidden_features, bias=True),
                nn.ReLU(inplace=True),
            ]
        )

    layers.extend(
        [nn.Linear(hidden_features, out_features, bias=True), nn.Sigmoid()]
    )

    self.net = nn.Sequential(*layers)

forward

forward(x)

Forward pass.

Parameters:

Name Type Description Default
x Tensor

Input tensor of shape (batch_size, in_features).

required

Returns:

Name Type Description
x Tensor

L1-normalized output tensor of shape (batch_size, out_features), summing to 1 along the last dim.

Source code in deeplens-src/deeplens/surrogate/mlp.py
def forward(self, x):
    """Forward pass.

    Args:
        x (torch.Tensor): Input tensor of shape `(batch_size, in_features)`.

    Returns:
        x (torch.Tensor): L1-normalized output tensor of shape
            `(batch_size, out_features)`, summing to 1 along the last dim.
    """
    x = self.net(x)
    x = F.normalize(x, p=1, dim=-1)
    return x

MLP with convolutional layers for spatial PSF prediction.

deeplens.surrogate.MLPConv

MLPConv(in_features, ks, channels=3, activation='relu')

Bases: Module

MLP encoder plus convolutional decoder for high-resolution PSF prediction.

The MLP encoder maps the input features to a low-resolution feature map of spatial size min(ks, 32), which a transposed-convolution decoder then upsamples by powers of two to the target PSF size ks. The decoder output is always Sigmoid-activated and L1-normalized over the two spatial dimensions so each predicted PSF sums to one.

Reference

"Differentiable Compound Optics and Processing Pipeline Optimization for End-To-end Camera Design".

Attributes:

Name Type Description
ks int

Spatial size of the output PSF.

ks_mlp int

Spatial size of the MLP feature map, min(ks, 32).

channels int

Number of output channels.

encoder Sequential

Linear encoder producing the feature map.

decoder Sequential

Transposed-convolution upsampling decoder.

activation Module

Activation module selected by activation. Note that the forward pass uses a Sigmoid regardless of this attribute.

Parameters:

Name Type Description Default
in_features int

Number of input features (e.g. field angle plus wavelength).

required
ks int

Spatial size of the output PSF. When greater than 32 it must be a multiple of 32 (asserted), and in practice \(32 \cdot 2^n\) so the decoder upsamples by integer powers of two.

required
channels int

Number of output channels. Defaults to 3.

3
activation str

Activation name, "relu" or "sigmoid", stored on self.activation but unused by forward. Defaults to "relu".

'relu'
Source code in deeplens-src/deeplens/surrogate/mlpconv.py
def __init__(self, in_features, ks, channels=3, activation="relu"):
    super(MLPConv, self).__init__()

    self.ks_mlp = min(ks, 32)
    upsample_times = 0  # ks <= 32 needs no upsampling (decoder loop runs 0 times)
    if ks > 32:
        assert ks % 32 == 0, "ks must be 32n"
        upsample_times = int(math.log(ks / 32, 2))

    linear_output = channels * self.ks_mlp**2
    self.ks = ks
    self.channels = channels

    # MLP encoder
    self.encoder = nn.Sequential(
        nn.Linear(in_features, 256),
        nn.ReLU(),
        nn.Linear(256, 256),
        nn.ReLU(),
        nn.Linear(256, 512),
        nn.ReLU(),
        nn.Linear(512, linear_output),
    )

    # Conv decoder
    conv_layers = []
    conv_layers.append(
        nn.ConvTranspose2d(channels, 64, kernel_size=3, stride=1, padding=1)
    )
    conv_layers.append(nn.ReLU())
    for _ in range(upsample_times):
        conv_layers.append(
            nn.ConvTranspose2d(64, 64, kernel_size=3, stride=1, padding=1)
        )
        conv_layers.append(nn.ReLU())
        conv_layers.append(nn.Upsample(scale_factor=2))

    conv_layers.append(
        nn.ConvTranspose2d(64, 64, kernel_size=3, stride=1, padding=1)
    )
    conv_layers.append(nn.ReLU())
    conv_layers.append(
        nn.ConvTranspose2d(64, channels, kernel_size=3, stride=1, padding=1)
    )
    self.decoder = nn.Sequential(*conv_layers)

    if activation == "relu":
        self.activation = nn.ReLU()
    elif activation == "sigmoid":
        self.activation = nn.Sigmoid()

forward

forward(x)

Predict normalized PSFs from input feature vectors.

Encodes x into a (batch_size, channels, ks_mlp, ks_mlp) feature map, upsamples it through the conv decoder, then applies Sigmoid and L1 normalization over the spatial dimensions so each PSF sums to one.

Parameters:

Name Type Description Default
x Tensor

Input tensor of shape (batch_size, in_features).

required

Returns:

Name Type Description
decoded Tensor

Normalized PSF tensor of shape (batch_size, channels, ks, ks).

Source code in deeplens-src/deeplens/surrogate/mlpconv.py
def forward(self, x):
    """Predict normalized PSFs from input feature vectors.

    Encodes `x` into a `(batch_size, channels, ks_mlp, ks_mlp)` feature map,
    upsamples it through the conv decoder, then applies Sigmoid and L1
    normalization over the spatial dimensions so each PSF sums to one.

    Args:
        x (torch.Tensor): Input tensor of shape `(batch_size, in_features)`.

    Returns:
        decoded (torch.Tensor): Normalized PSF tensor of shape
            `(batch_size, channels, ks, ks)`.
    """
    # Encode the input using the MLP
    encoded = self.encoder(x)

    # Reshape the output from the MLP to feed to the CNN
    decoded_input = encoded.view(
        -1, self.channels, self.ks_mlp, self.ks_mlp
    )  # reshape to (batch_size, channels, height, width)

    # Decode the output using the CNN
    decoded = self.decoder(decoded_input)

    # This normalization only works for PSF network
    decoded = nn.Sigmoid()(decoded)
    decoded = F.normalize(decoded, p=1, dim=[-1, -2])

    return decoded

Sinusoidal-activation network (SIREN) for representing high-frequency PSF detail.

deeplens.surrogate.siren.Siren

Siren(dim_in, dim_out, w0=1.0, c=6.0, is_first=False, use_bias=True, activation=None)

Bases: Module

Single SIREN (Sinusoidal Representation Network) layer.

A linear layer followed by a sine activation. Uses the initialization scheme from "Implicit Neural Representations with Periodic Activation Functions".

Parameters:

Name Type Description Default
dim_in int

Input dimension.

required
dim_out int

Output dimension.

required
w0 float

Frequency multiplier for the sine activation. Defaults to 1.0.

1.0
c float

Constant controlling the weight initialization scale (non-first layers). Defaults to 6.0.

6.0
is_first bool

Whether this is the first layer (uses a different init scale). Defaults to False.

False
use_bias bool

Whether to include a bias term. Defaults to True.

True
activation Module or None

Custom activation module. Defaults to None, which uses Sine(w0).

None
Source code in deeplens-src/deeplens/surrogate/siren.py
def __init__(
    self,
    dim_in,
    dim_out,
    w0=1.0,
    c=6.0,
    is_first=False,
    use_bias=True,
    activation=None,
):
    super().__init__()
    self.dim_in = dim_in
    self.is_first = is_first

    weight = torch.zeros(dim_out, dim_in)
    bias = torch.zeros(dim_out) if use_bias else None
    self.init_(weight, bias, c=c, w0=w0)

    self.weight = nn.Parameter(weight)
    self.bias = nn.Parameter(bias) if use_bias else None
    self.activation = Sine(w0) if activation is None else activation

init_

init_(weight, bias, c, w0)

Initialize the layer weight in place with the SIREN scheme.

Fills weight uniformly in \([-w_{std}, w_{std}]\), where the std is \(1/\text{dim}\) for the first layer and \(\sqrt{c/\text{dim}}/w_0\) otherwise. The bias argument is accepted for API symmetry but left unchanged (it stays at its zero-initialized value).

Parameters:

Name Type Description Default
weight Tensor

Weight tensor of shape (dim_out, dim_in), modified in place.

required
bias Tensor or None

Bias tensor of shape (dim_out,), not modified.

required
c float

Constant controlling the initialization scale.

required
w0 float

Frequency multiplier for the sine activation.

required
Source code in deeplens-src/deeplens/surrogate/siren.py
def init_(self, weight, bias, c, w0):
    """Initialize the layer weight in place with the SIREN scheme.

    Fills `weight` uniformly in $[-w_{std}, w_{std}]$, where the std is
    $1/\\text{dim}$ for the first layer and $\\sqrt{c/\\text{dim}}/w_0$
    otherwise. The `bias` argument is accepted for API symmetry but left
    unchanged (it stays at its zero-initialized value).

    Args:
        weight (torch.Tensor): Weight tensor of shape `(dim_out, dim_in)`, modified in place.
        bias (torch.Tensor or None): Bias tensor of shape `(dim_out,)`, not modified.
        c (float): Constant controlling the initialization scale.
        w0 (float): Frequency multiplier for the sine activation.
    """
    dim = self.dim_in

    w_std = (1 / dim) if self.is_first else (math.sqrt(c / dim) / w0)
    weight.uniform_(-w_std, w_std)

forward

forward(x)

Forward pass.

Parameters:

Name Type Description Default
x Tensor

Input tensor of shape (..., dim_in).

required

Returns:

Name Type Description
out Tensor

Output tensor of shape (..., dim_out).

Source code in deeplens-src/deeplens/surrogate/siren.py
def forward(self, x):
    """Forward pass.

    Args:
        x (torch.Tensor): Input tensor of shape `(..., dim_in)`.

    Returns:
        out (torch.Tensor): Output tensor of shape `(..., dim_out)`.
    """
    out = F.linear(x, self.weight, self.bias)
    out = self.activation(out)
    return out

SIREN variant with feature modulation for conditioning on lens parameters.

deeplens.surrogate.ModulateSiren

ModulateSiren(
    dim_in,
    dim_hidden,
    dim_out,
    dim_latent,
    num_layers,
    image_width,
    image_height,
    w0=1.0,
    w0_initial=30.0,
    use_bias=True,
    final_activation=None,
    outermost_linear=True,
)

Bases: Module

Modulated SIREN for latent-conditioned image synthesis.

Combines a SIREN synthesizer network (mapping a fixed pixel-coordinate grid to output values) with a modulator network that scales each synthesizer layer based on a conditioning latent vector. Used to predict spatially-varying PSFs conditioned on lens parameters. The output is always tanh-activated and reshaped to an image regardless of the outermost_linear / final_activation settings.

Attributes:

Name Type Description
synthesizer ModuleList

SIREN sine layers plus the final output layer.

modulator ModuleList

Per-layer Linear+ReLU blocks producing modulation vectors from the latent (and previous modulation).

grid Tensor

Registered coordinate buffer of shape (image_height * image_width, dim_in), spanning \([-1, 1]\) on each axis.

Parameters:

Name Type Description Default
dim_in int

Input coordinate dimension (typically 2 for x, y).

required
dim_hidden int

Hidden layer width for both synthesizer and modulator.

required
dim_out int

Output dimension per pixel (e.g., 1 for grayscale PSF).

required
dim_latent int

Dimension of the conditioning latent vector.

required
num_layers int

Number of SIREN + modulator layers (excluding the final output layer of the synthesizer).

required
image_width int

Output image width in pixels.

required
image_height int

Output image height in pixels.

required
w0 float

Frequency multiplier for hidden sine layers. Defaults to 1.0.

1.0
w0_initial float

Frequency multiplier for the first sine layer. Defaults to 30.0.

30.0
use_bias bool

Whether to use bias in sine layers. Defaults to True.

True
final_activation Module or None

Activation for the final Siren layer when outermost_linear is False. Defaults to None (Identity).

None
outermost_linear bool

If True, the final synthesizer layer is a plain nn.Linear; otherwise it is a Siren layer. Defaults to True.

True
Source code in deeplens-src/deeplens/surrogate/modulate_siren.py
def __init__(
    self,
    dim_in,
    dim_hidden,
    dim_out,
    dim_latent,
    num_layers,
    image_width,
    image_height,
    w0=1.0,
    w0_initial=30.0,
    use_bias=True,
    final_activation=None,
    outermost_linear=True,
):
    super().__init__()
    self.num_layers = num_layers
    self.dim_hidden = dim_hidden
    self.img_width = image_width
    self.img_height = image_height

    # ==> Synthesizer
    synthesizer_layers = nn.ModuleList([])
    for ind in range(num_layers):
        is_first = ind == 0
        layer_w0 = w0_initial if is_first else w0
        layer_dim_in = dim_in if is_first else dim_hidden

        synthesizer_layers.append(
            SineLayer(
                in_features=layer_dim_in,
                out_features=dim_hidden,
                omega_0=layer_w0,
                bias=use_bias,
                is_first=is_first,
            )
        )

    if outermost_linear:
        last_layer = nn.Linear(dim_hidden, dim_out)
        with torch.no_grad():
            # w_std = math.sqrt(6 / dim_hidden) / w0
            # self.last_layer.weight.uniform_(- w_std, w_std)
            nn.init.kaiming_normal_(
                last_layer.weight, a=0.0, nonlinearity="relu", mode="fan_in"
            )
    else:
        final_activation = (
            nn.Identity() if not exists(final_activation) else final_activation
        )
        last_layer = Siren(
            dim_in=dim_hidden,
            dim_out=dim_out,
            w0=w0,
            use_bias=use_bias,
            activation=final_activation,
        )
    synthesizer_layers.append(last_layer)

    self.synthesizer = synthesizer_layers
    # self.synthesizer = nn.Sequential(*synthesizer)

    # ==> Modulator
    modulator_layers = nn.ModuleList([])
    for ind in range(num_layers):
        is_first = ind == 0
        dim = dim_latent if is_first else (dim_hidden + dim_latent)

        modulator_layers.append(
            nn.Sequential(nn.Linear(dim, dim_hidden), nn.ReLU())
        )

        with torch.no_grad():
            # self.layers[-1][0].weight.uniform_(-1 / dim_hidden, 1 / dim_hidden)
            nn.init.kaiming_normal_(
                modulator_layers[-1][0].weight,
                a=0.0,
                nonlinearity="relu",
                mode="fan_in",
            )

    self.modulator = modulator_layers
    # self.modulator = nn.Sequential(*modulator_layers)

    # ==> Positions
    tensors = [
        torch.linspace(-1, 1, steps=image_height),
        torch.linspace(-1, 1, steps=image_width),
    ]
    mgrid = torch.stack(torch.meshgrid(*tensors, indexing="ij"), dim=-1)
    mgrid = rearrange(mgrid, "h w c -> (h w) c")
    self.register_buffer("grid", mgrid)

forward

forward(latent)

Synthesize a batch of images from conditioning latent vectors.

Runs the shared coordinate grid through the SIREN synthesizer, scaling each layer by the corresponding modulator output, then applies a tanh and reshapes to a channel-first image batch.

Parameters:

Name Type Description Default
latent Tensor

Conditioning latent vector of shape (batch_size, dim_latent).

required

Returns:

Name Type Description
x Tensor

Output image tensor of shape (batch_size, 1, image_height, image_width), with values in \([-1, 1]\).

Source code in deeplens-src/deeplens/surrogate/modulate_siren.py
def forward(self, latent):
    """Synthesize a batch of images from conditioning latent vectors.

    Runs the shared coordinate grid through the SIREN synthesizer, scaling each
    layer by the corresponding modulator output, then applies a tanh and reshapes
    to a channel-first image batch.

    Args:
        latent (torch.Tensor): Conditioning latent vector of shape
            `(batch_size, dim_latent)`.

    Returns:
        x (torch.Tensor): Output image tensor of shape
            `(batch_size, 1, image_height, image_width)`, with values in
            $[-1, 1]$.
    """
    x = self.grid.clone().detach().requires_grad_()

    for i in range(self.num_layers):
        if i == 0:
            z = self.modulator[i](latent)
        else:
            z = self.modulator[i](torch.cat((latent, z), dim=-1))

        x = self.synthesizer[i](x)
        x = x * z

    x = self.synthesizer[-1](x)  # shape of (h*w, 1)
    x = torch.tanh(x)
    x = x.view(
        -1, self.img_height, self.img_width, 1
    )  # reshape to (batch_size, height, width, channels)
    x = x.permute(0, 3, 1, 2)  # reshape to (batch_size, channels, height, width)
    return x