cuwave 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
cuwave/maxwell.py ADDED
@@ -0,0 +1,416 @@
1
+ """Maxwell's equations in second-order curl-curl form, and their scalar 2D reductions.
2
+
3
+ `MaxwellWave` is the Yee lattice written as `-C^T nu C` with a diagonal permittivity,
4
+ which is the staggered `ElasticWave` layout with its normal-stress block deleted, and
5
+ it is what 3D needs. Out of plane one component survives and the curl-curl collapses to
6
+ a flux divergence, so 2D needs neither: `ElectricWave` carries E_z with the permittivity
7
+ as its inertia and `MagneticWave` carries H_z with the inverse permittivity as its
8
+ stiffness, both `PressureWave` on the scalar kernels.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import math
14
+ from collections.abc import Callable
15
+ from dataclasses import dataclass
16
+ from pathlib import Path
17
+
18
+ import cupy as cp
19
+ import cupy.typing as cpt
20
+ import numpy as np
21
+ import numpy.typing as npt
22
+
23
+ from .boundary import Conductor, Magnetic, faces_with
24
+ from .scalar import PressureWave
25
+ from .wave import (
26
+ PAIRS,
27
+ Simulation,
28
+ apply_cell_weights,
29
+ axis_geometry,
30
+ component_weights,
31
+ grid_block,
32
+ pair_average,
33
+ pair_average_adjoint,
34
+ pair_weights,
35
+ point_average,
36
+ point_average_adjoint,
37
+ )
38
+
39
+ KERNEL_PATH = Path(__file__).parent / "kernels" / "maxwell.cu"
40
+ SENSITIVITY_PATH = Path(__file__).parent / "kernels" / "maxwell_sensitivity.cu"
41
+
42
+
43
+ # -------------------------------------- helpers --------------------------------------
44
+ def interpolate(indicator: cpt.NDArray, first: float, second: float) -> cpt.NDArray:
45
+ """Two-phase linear interpolation `first + gamma (second - first)` of a material."""
46
+ return first + indicator * (second - first)
47
+
48
+
49
+ # -------------------------------- discretization setup -------------------------------
50
+ @dataclass
51
+ class PolarizedWave(PressureWave):
52
+ """Two-phase 2D Maxwell base: one out-of-plane component, gamma interpolating epsilon."""
53
+
54
+ permittivity1: float = None # background, gamma = 0
55
+ permittivity2: float = None # design material, gamma = 1
56
+ permeability: float = 1.0 # the media are non-magnetic, so it is not designed
57
+
58
+ def __post_init__(self) -> None:
59
+ """Validate the two phases and the out-of-plane reduction, on top of `Simulation`."""
60
+ super().__post_init__()
61
+ if None in (self.permittivity1, self.permittivity2):
62
+ raise ValueError("PolarizedWave requires permittivity1 and permittivity2")
63
+ if min(self.permittivity1, self.permittivity2, self.permeability) <= 0.0:
64
+ raise ValueError(
65
+ f"permittivity and permeability must be positive: "
66
+ f"{self.permittivity1}, {self.permittivity2}, {self.permeability}"
67
+ )
68
+ if self.ndim > 2:
69
+ raise ValueError(f"an out-of-plane reduction is 1D or 2D, not {self.ndim}D")
70
+
71
+ @property
72
+ def light_speed(self) -> float:
73
+ """Fastest speed on the grid, `1 / sqrt(min(eps) mu)`: what bounds the timestep."""
74
+ # the background is the fast phase here, the opposite of the seismic case
75
+ return 1.0 / math.sqrt(
76
+ min(self.permittivity1, self.permittivity2) * self.permeability
77
+ )
78
+
79
+ def permittivity(self, indicator: cpt.NDArray) -> cpt.NDArray:
80
+ """Two-phase permittivity `eps1 + gamma (eps2 - eps1)`, linear in the design."""
81
+ return interpolate(indicator, self.permittivity1, self.permittivity2)
82
+
83
+ def constant(self, value: float) -> cpt.NDArray:
84
+ """`value` over the padded grid, for the coefficient the design does not enter."""
85
+ return cp.full(self.Nx_padded, value, dtype=self.dtype)
86
+
87
+ def step_factors(self) -> list:
88
+ """Per-axis finite-difference step factors `2 * dt**2 / dx**2`."""
89
+ return [self.dtype(2.0 * self.dt**2 / dxk**2) for dxk in self.dx]
90
+
91
+ def source_factor(self) -> float:
92
+ """Source scaling, unscaled since both materials already sit in the fields."""
93
+ return 1.0
94
+
95
+
96
+ @dataclass
97
+ class ElectricWave(PolarizedWave):
98
+ """Out-of-plane electric field, the permittivity its inertia and `1 / mu` its stiffness.
99
+
100
+ The polarization Christiansen & Sigmund 2021 (https://doi.org/10.1364/JOSAB.406048)
101
+ label TE, and the one a waveguide device is designed in, since a guided mode's
102
+ electric field is the one held continuous across the sidewalls.
103
+ """
104
+
105
+ def parametrization(
106
+ self, indicator: cpt.NDArray
107
+ ) -> tuple[cpt.NDArray, cpt.NDArray]:
108
+ """`indicator` interpolates the permittivity, which is the inertia of E_z."""
109
+ nu = self.constant(1.0 / self.permeability)
110
+ return nu, 1.0 / self.permittivity(indicator)
111
+
112
+ def parametrization_jacobian(self, indicator: cpt.NDArray) -> tuple:
113
+ """The inertia is affine in gamma and the stiffness does not depend on it at all."""
114
+ return self.permittivity2 - self.permittivity1, 0.0
115
+
116
+
117
+ @dataclass
118
+ class MagneticWave(PolarizedWave):
119
+ """Out-of-plane magnetic field, `mu` its inertia and the inverse permittivity its stiffness.
120
+
121
+ The polarization Christiansen & Sigmund 2021 label TM, and the one their metalens
122
+ is designed in. The design enters the stiffness as `1 / eps`, so its jacobian is a
123
+ field rather than the constant `ElectricWave` gets.
124
+ """
125
+
126
+ def parametrization(
127
+ self, indicator: cpt.NDArray
128
+ ) -> tuple[cpt.NDArray, cpt.NDArray]:
129
+ """`indicator` interpolates the permittivity, which enters H_z as `1 / eps`."""
130
+ nu = self.constant(1.0 / self.permeability)
131
+ return 1.0 / self.permittivity(indicator), nu
132
+
133
+ def parametrization_jacobian(self, indicator: cpt.NDArray) -> tuple:
134
+ """The inertia is constant; `1 / eps` is not affine in gamma, so its derivative is nodal."""
135
+ contrast = self.permittivity2 - self.permittivity1
136
+ return 0.0, -contrast / self.permittivity(indicator) ** 2
137
+
138
+
139
+ # ---------------------------------- the vector case ----------------------------------
140
+ @dataclass
141
+ class MaxwellWave(Simulation):
142
+ """Curl-curl Maxwell on the Yee lattice, the staggered scheme 3D needs.
143
+
144
+ Component `c` of the electric field lives half a node up axis `c` and the curl of
145
+ the pair `(k, l)` half a node up both of its axes, which is the `ElasticWave`
146
+ displacement and shear-stress layout exactly. The operator is assembled as the
147
+ variational derivative of the magnetic energy, so it is `-C^T nu C` with the
148
+ Levi-Civita signs squared away, symmetric by construction and the exact transpose
149
+ at every order.
150
+ """
151
+
152
+ permeability: float = 1.0 # background mu, designed only where `magnetic` is set
153
+
154
+ kernel_path = KERNEL_PATH
155
+ sensitivity_path = SENSITIVITY_PATH
156
+ default_boundary = Conductor
157
+ gradient_names = ("mass", "stiff")
158
+
159
+ magnetic = False # set where the permeability carries the design too
160
+
161
+ @property
162
+ def compile_flags(self) -> tuple[str, ...]:
163
+ """`-DUSE_MAGNETIC` on top of the damping flag, where `mu` carries the design too."""
164
+ flags = super().compile_flags
165
+ return (*flags, "-DUSE_MAGNETIC") if self.magnetic else flags
166
+
167
+ @property
168
+ def ncomp(self) -> int:
169
+ """One electric field component per axis."""
170
+ return self.ndim
171
+
172
+ @property
173
+ def component_offsets(self) -> npt.NDArray[np.float64]:
174
+ """Component `c` sits half a node up axis `c`: the Yee staggering itself."""
175
+ return 0.5 * np.eye(self.ndim)
176
+
177
+ @property
178
+ def npairs(self) -> int:
179
+ """Curl components, one per axis pair: 1 in 2D and 3 in 3D."""
180
+ return len(PAIRS[self.ndim]) - self.ndim
181
+
182
+ @property
183
+ def radius(self) -> int:
184
+ """Half the stencil width, so `space_order` is twice it."""
185
+ return self.space_order // 2
186
+
187
+ @property
188
+ def reach(self) -> int:
189
+ """The curl reach compounds with the transposed one: `2 radius - 1` nodes."""
190
+ return 2 * self.radius - 1
191
+
192
+ def __post_init__(self) -> None:
193
+ """Validate the dimension and the faces, on top of `Simulation.__post_init__`."""
194
+ super().__post_init__()
195
+ if self.ndim < 2:
196
+ raise ValueError(
197
+ "a single in-plane component has no curl; use ElectricWave"
198
+ )
199
+ if self.permeability <= 0.0:
200
+ raise ValueError(f"permeability must be positive: {self.permeability}")
201
+ for pair in self.boundary:
202
+ for condition in pair:
203
+ if condition not in (Conductor, Magnetic):
204
+ raise ValueError(f"maxwell faces are Conductor or Magnetic: {pair}")
205
+
206
+ def inverse_inertia(self, indicator: cpt.NDArray) -> cpt.NDArray:
207
+ """Nodal `1 / (eps W)`, what a sponge scales its conductivity by."""
208
+ permittivity, _ = self.parametrization(indicator)
209
+ mass = permittivity * apply_cell_weights(
210
+ self, cp.ones(self.Nx_padded, dtype=self.dtype)
211
+ )
212
+ return 1.0 / cp.maximum(mass, cp.finfo(self.dtype).tiny)
213
+
214
+ def build_materials(self, indicator: cpt.NDArray) -> dict:
215
+ """Point inverse permittivity per component, and the pair inverse permeability."""
216
+ permittivity, nu = self.parametrization(indicator)
217
+ permittivity = cp.ascontiguousarray(permittivity, dtype=self.dtype)
218
+ minv = cp.zeros((self.ncomp, *self.Nx_padded), dtype=self.dtype)
219
+ for c in range(self.ncomp):
220
+ mass = component_weights(self, c) * point_average(self, permittivity, c)
221
+ minv[c] = 1.0 / cp.maximum(mass, cp.finfo(self.dtype).tiny)
222
+ # a conductor holds the tangential field, which is what a zeroed inertia does
223
+ for face in faces_with(self, Conductor):
224
+ wall = [slice(None)] * self.ndim
225
+ wall[face // 2] = 1 if face % 2 == 0 else self.Nx[face // 2] - 2
226
+ for c in range(self.ncomp):
227
+ if c != face // 2:
228
+ minv[(c, *wall)] = 0.0
229
+ mat = {"minv": cp.ascontiguousarray(minv), "eps": permittivity}
230
+ if self.magnetic:
231
+ nu = cp.ascontiguousarray(nu, dtype=self.dtype)
232
+ pairs = cp.zeros((self.npairs, *self.Nx_padded), dtype=self.dtype)
233
+ for p, axes in enumerate(PAIRS[self.ndim][self.ndim :]):
234
+ pairs[p] = pair_weights(self, axes) * pair_average(self, nu, axes)
235
+ mat["nu_pair"] = cp.ascontiguousarray(pairs)
236
+ mat["nu"] = nu
237
+ if self.damping is not None:
238
+ mat["damping"] = self.damping
239
+ return mat
240
+
241
+ def define_step(self, kernels: cp.RawModule, mat: dict) -> Callable:
242
+ """Closure launching the curl kernel and then the update over (u0, u1, u2)."""
243
+ curl_kernel = kernels.get_function("curl_kernel")
244
+ fd_kernel = kernels.get_function("fd_kernel")
245
+ grid, block = grid_block(self)
246
+ # the curl scratch outlives the closure, its invalid pair points never written
247
+ h = cp.zeros((self.npairs, *self.Nx_padded), dtype=self.dtype)
248
+ inv_dx = [self.dtype(1.0 / d) for d in self.dx]
249
+ cargs = [None, h]
250
+ if self.magnetic:
251
+ cargs.append(mat["nu_pair"])
252
+ cargs += [
253
+ self.dtype(1.0 / self.permeability),
254
+ np.int32(self.comp_stride),
255
+ *axis_geometry(self, inv_dx),
256
+ ]
257
+ uargs = [None, None, None, h, mat["minv"]]
258
+ if self.damping is not None:
259
+ uargs += [mat["damping"], self.dtype(self.dt)]
260
+ uargs += [np.int32(self.comp_stride), *axis_geometry(self, self.step_factors())]
261
+
262
+ def fd_step(u0, u1, u2):
263
+ cargs[0] = u1
264
+ curl_kernel(grid, block, cargs)
265
+ uargs[0], uargs[1], uargs[2] = u0, u1, u2
266
+ fd_kernel(grid, block, uargs)
267
+ return u2
268
+
269
+ return fd_step
270
+
271
+ def excitation_weights(
272
+ self, mat: dict, lin_index: cpt.NDArray[cp.int32]
273
+ ) -> cpt.NDArray:
274
+ """Source weights `dt**2 / (eps V)`: the kernel inertia leaves the volume out."""
275
+ volume = self.dtype(1.0 / float(np.prod(self.dx)))
276
+ weight = mat["minv"].ravel()[lin_index] * volume
277
+ if self.damping is not None:
278
+ node = lin_index % np.int32(self.comp_stride)
279
+ beta = 0.5 * weight * mat["damping"].ravel()[node] * self.dt
280
+ weight = weight / (1.0 + beta)
281
+ return (self.dtype(self.dt**2 * self.source_factor()) * weight).astype(
282
+ self.dtype
283
+ )
284
+
285
+ def adjoint_weights(self, sensors: cpt.NDArray[cp.int32]) -> cpt.NDArray:
286
+ """`1 / V`: dJ/du is nodal, so it undoes the volume the source weights divide by."""
287
+ volume = self.dtype(1.0 / float(np.prod(self.dx)))
288
+ return cp.full(sensors.shape[1], volume, dtype=self.dtype)
289
+
290
+ def step_factors(self) -> list:
291
+ """Per-axis update factors `dt**2 / dx`, the curl carrying the other `1 / dx`."""
292
+ return [self.dtype(self.dt**2 / d) for d in self.dx]
293
+
294
+ def source_factor(self) -> float:
295
+ """Source scaling, unscaled since the permittivity is already in the point inertia."""
296
+ return 1.0
297
+
298
+ def gradient_fields(self, mat: dict) -> dict[str, cpt.NDArray]:
299
+ """Accumulators on the component points, and on the pair points where `mu` is designed."""
300
+ grads = {
301
+ "mass": cp.zeros((self.ncomp, *self.Nx_padded), dtype=self.dtype),
302
+ }
303
+ if self.magnetic:
304
+ grads["nu"] = cp.zeros((self.npairs, *self.Nx_padded), dtype=self.dtype)
305
+ # the harmonic mean has a design dependent chain rule, so keep the field
306
+ grads["design"] = mat["nu"]
307
+ return grads
308
+
309
+ def finalize_gradients(self, grads: dict, kernels: cp.RawModule) -> dict:
310
+ """Chain the point densities through the averages onto the nodal design fields."""
311
+ g_mass = cp.zeros(self.Nx_padded, dtype=self.dtype)
312
+ for c in range(self.ncomp):
313
+ density = grads["mass"][c] * component_weights(self, c)
314
+ g_mass += point_average_adjoint(self, density, c)
315
+ # a non-magnetic medium has no stiffness design dependence: zeros say so
316
+ g_stiff = cp.zeros(self.Nx_padded, dtype=self.dtype)
317
+ if self.magnetic:
318
+ for p, axes in enumerate(PAIRS[self.ndim][self.ndim :]):
319
+ density = grads["nu"][p] * pair_weights(self, axes)
320
+ g_stiff += pair_average_adjoint(self, density, grads["design"], axes)
321
+ return {"mass": g_mass, "stiff": g_stiff}
322
+
323
+ def _density_args(self, grads: dict) -> list:
324
+ """The accumulator head shared by the gradient and Frechet closures."""
325
+ args = [grads["mass"]]
326
+ if self.magnetic:
327
+ args.append(grads["nu"])
328
+ return args
329
+
330
+ def define_gradient(
331
+ self, kernels: cp.RawModule, mat: dict, grads: dict
332
+ ) -> Callable:
333
+ """Closure accumulating the point permittivity density from a triplet and `l1`."""
334
+ gradient_kernel = kernels.get_function("gradient_kernel")
335
+ grid, block = grid_block(self)
336
+ inv_dx = [self.dtype(1.0 / d) for d in self.dx]
337
+ head = self._density_args(grads)
338
+ args = (
339
+ head
340
+ + [None, None, None, None]
341
+ + [
342
+ self.dtype(1.0 / self.dt**2),
343
+ np.int32(self.comp_stride),
344
+ *axis_geometry(self, inv_dx),
345
+ ]
346
+ )
347
+ start = len(head)
348
+
349
+ def gradient_step(u0, u1, u2, l1):
350
+ args[start], args[start + 1] = u0, u1
351
+ args[start + 2], args[start + 3] = u2, l1
352
+ gradient_kernel(grid, block, args)
353
+
354
+ return gradient_step
355
+
356
+ def define_frechet(
357
+ self, kernels: cp.RawModule, accs: dict, sign: float
358
+ ) -> Callable:
359
+ """Closure accumulating the quadratic densities of one field triplet, times `sign`."""
360
+ frechet_kernel = kernels.get_function("frechet_kernel")
361
+ grid, block = grid_block(self)
362
+ inv_dx = [self.dtype(1.0 / d) for d in self.dx]
363
+ head = self._density_args(accs)
364
+ args = (
365
+ head
366
+ + [None, None, None]
367
+ + [
368
+ self.dtype(sign / (2.0 * self.dt) ** 2),
369
+ self.dtype(-sign),
370
+ np.int32(self.comp_stride),
371
+ *axis_geometry(self, inv_dx),
372
+ ]
373
+ )
374
+ start = len(head)
375
+
376
+ def frechet_step(u0, u1, u2):
377
+ args[start], args[start + 1], args[start + 2] = u0, u1, u2
378
+ frechet_kernel(grid, block, args)
379
+
380
+ return frechet_step
381
+
382
+
383
+ @dataclass
384
+ class DielectricWave(MaxwellWave):
385
+ """Two-phase dielectric, gamma interpolating the permittivity between the two phases."""
386
+
387
+ permittivity1: float = None # background, gamma = 0
388
+ permittivity2: float = None # design material, gamma = 1
389
+
390
+ def __post_init__(self) -> None:
391
+ """Validate the two phases, on top of `MaxwellWave.__post_init__`."""
392
+ super().__post_init__()
393
+ if None in (self.permittivity1, self.permittivity2):
394
+ raise ValueError("DielectricWave requires permittivity1 and permittivity2")
395
+ if min(self.permittivity1, self.permittivity2) <= 0.0:
396
+ raise ValueError(
397
+ f"permittivity must be positive: "
398
+ f"{self.permittivity1}, {self.permittivity2}"
399
+ )
400
+
401
+ @property
402
+ def light_speed(self) -> float:
403
+ """Fastest speed on the grid, `1 / sqrt(min(eps) mu)`: what bounds the timestep."""
404
+ return 1.0 / math.sqrt(
405
+ min(self.permittivity1, self.permittivity2) * self.permeability
406
+ )
407
+
408
+ def parametrization(
409
+ self, indicator: cpt.NDArray
410
+ ) -> tuple[cpt.NDArray, cpt.NDArray | None]:
411
+ """`indicator` interpolates the permittivity; the permeability is left uniform."""
412
+ return interpolate(indicator, self.permittivity1, self.permittivity2), None
413
+
414
+ def parametrization_jacobian(self, indicator: cpt.NDArray) -> tuple:
415
+ """The permittivity is affine in gamma, and the permeability does not depend on it."""
416
+ return self.permittivity2 - self.permittivity1, 0.0
cuwave/nn.py ADDED
@@ -0,0 +1,99 @@
1
+ """Neural networks as a reparametrization of the design field, in PyTorch.
2
+
3
+ See Herrmann, Buerchner, Dietrich & Kollmannsberger 2023
4
+ (https://doi.org/10.1016/j.cma.2023.116278) and Herrmann, Sigmund, Li, Vogl &
5
+ Kollmannsberger 2024 (https://doi.org/10.1007/s00158-024-03908-6).
6
+ """
7
+
8
+ import torch
9
+ from torch import nn
10
+
11
+ CONVOLUTIONS = {1: nn.Conv1d, 2: nn.Conv2d, 3: nn.Conv3d}
12
+ OUTPUT_STD = 0.01
13
+
14
+
15
+ # -------------------------------------- helpers --------------------------------------
16
+ def _initialize(model: nn.Module) -> None:
17
+ """Xavier uniform with zero bias."""
18
+ for module in model.modules():
19
+ if isinstance(module, tuple(CONVOLUTIONS.values())):
20
+ nn.init.xavier_uniform_(module.weight)
21
+ nn.init.zeros_(module.bias)
22
+
23
+
24
+ def nn_params(model: nn.Module) -> int:
25
+ """Number of trainable parameters in `model`"""
26
+ return sum(p.numel() for p in model.parameters() if p.requires_grad)
27
+
28
+
29
+ def _convolution(
30
+ in_channels: int, out_channels: int, kernel: int, dim: int
31
+ ) -> nn.Module:
32
+ """`kernel`-wide convolution over `dim` axes, padded to retain the resolution"""
33
+ if dim not in CONVOLUTIONS:
34
+ raise ValueError(f"dim is 1, 2 or 3, not {dim!r}")
35
+ return CONVOLUTIONS[dim](in_channels, out_channels, kernel, padding=kernel // 2)
36
+
37
+
38
+ # -------------------------------------- networks -------------------------------------
39
+ class Generator(nn.Module):
40
+ """Fixed noise to field, doubling the resolution at every channel transition.
41
+
42
+ Nearest-neighbour upsampling per transition, and the stack ends in a sigmoid, so
43
+ the field is in [0, 1]. The output convolution starts with small random weights and
44
+ `output_bias`, so the field starts flat at 1, where the pixel-wise drivers start.
45
+
46
+ Args:
47
+ channels: latent channels, tapering to the one channel of the field.
48
+ shape: field resolution, divisible by 2**(len(channels) - 1).
49
+ kernel: convolution width, padded to retain the resolution.
50
+ activation: hidden activation; the output activation is always a sigmoid.
51
+ dim: 1, 2 or 3.
52
+ output_bias: bias into the output sigmoid, so a large value starts the field
53
+ flat at 1.
54
+ learnable: make `latent` a parameter instead of a buffer.
55
+ """
56
+
57
+ def __init__(
58
+ self,
59
+ channels: list[int],
60
+ shape: tuple[int, ...],
61
+ kernel: int = 5,
62
+ activation: type[nn.Module] = nn.GELU,
63
+ dim: int = 2,
64
+ output_bias: float = 10.0,
65
+ learnable: bool = False,
66
+ ) -> None:
67
+ super().__init__()
68
+ if len(shape) != dim:
69
+ raise ValueError(f"a {dim}D generator needs {dim} sizes, got {shape}")
70
+ blocks = len(channels) - 1
71
+ step = 2**blocks
72
+ if any(n % step for n in shape):
73
+ raise ValueError(f"{blocks} upsamplings need {shape} divisible by {step}")
74
+
75
+ layers = []
76
+ for i in range(blocks):
77
+ last = i == blocks - 1
78
+ layers += [
79
+ nn.Upsample(scale_factor=2, mode="nearest"),
80
+ nn.GroupNorm(1, channels[i]),
81
+ _convolution(channels[i], channels[i + 1], kernel, dim),
82
+ nn.Sigmoid() if last else activation(),
83
+ ]
84
+ self.stack = nn.Sequential(*layers)
85
+ _initialize(self)
86
+
87
+ output = self.stack[-2] # the convolution the sigmoid reads
88
+ nn.init.normal_(output.weight, std=OUTPUT_STD)
89
+ nn.init.constant_(output.bias, output_bias)
90
+
91
+ noise = torch.randn(1, channels[0], *(n // step for n in shape))
92
+ noise = 2.0 * noise / (noise.max() - noise.min()) # span activation's range
93
+ if learnable:
94
+ self.latent = nn.Parameter(noise)
95
+ else:
96
+ self.register_buffer("latent", noise)
97
+
98
+ def forward(self) -> torch.Tensor:
99
+ return self.stack(self.latent)
cuwave/optimization.py ADDED
@@ -0,0 +1,123 @@
1
+ from collections import deque
2
+ from collections.abc import Callable
3
+
4
+ import cupy.typing as cpt
5
+
6
+
7
+ class Adam:
8
+ """Adam optimizer (Kingma & Ba 2015, https://doi.org/10.48550/arXiv.1412.6980)"""
9
+
10
+ def __init__(
11
+ self,
12
+ lr: float = 1e-2,
13
+ betas: tuple[float, float] = (0.9, 0.999),
14
+ eps: float = 1e-8,
15
+ ) -> None:
16
+ self.lr = lr
17
+ self.b1, self.b2 = betas
18
+ self.eps = eps
19
+ self.m = self.v = 0.0
20
+ self.t = 0
21
+
22
+ def step(self, x: cpt.NDArray, grad: cpt.NDArray) -> cpt.NDArray:
23
+ """Update `x` using the Adam rule for one gradient `grad`."""
24
+ self.t += 1
25
+ self.m = self.b1 * self.m + (1.0 - self.b1) * grad
26
+ self.v = self.b2 * self.v + (1.0 - self.b2) * grad**2
27
+ m_hat = self.m / (1.0 - self.b1**self.t)
28
+ v_hat = self.v / (1.0 - self.b2**self.t)
29
+ return x - self.lr * m_hat / (v_hat**0.5 + self.eps)
30
+
31
+
32
+ class Lbfgs:
33
+ """L-BFGS with Armijo line search over the last `k` secant pairs.
34
+ Adapted from Fichtner 2021 (https://doi.org/10.33774/coe-2021-qpq2j).
35
+ """
36
+
37
+ def __init__(
38
+ self,
39
+ k: int = 10,
40
+ lr: float = 1.0,
41
+ first_step: float = 0.05,
42
+ armijo: float = 1e-4,
43
+ shrink: float = 0.5,
44
+ min_alpha: float = 1e-3,
45
+ ) -> None:
46
+ self.lr = lr
47
+ self.pairs = deque(maxlen=k) # oldest first, drops the oldest when full
48
+ self.x = self.grad = None
49
+ self.first_step = first_step
50
+ self.armijo, self.shrink, self.min_alpha = armijo, shrink, min_alpha
51
+
52
+ def search(
53
+ self,
54
+ x: cpt.NDArray,
55
+ grad: cpt.NDArray,
56
+ cost: float,
57
+ f: Callable,
58
+ project: Callable = lambda x: x,
59
+ ) -> tuple[cpt.NDArray, float, float, int]:
60
+ """Armijo backtracking on the proposal `step(x, grad)`.
61
+
62
+ Args:
63
+ x: the current design.
64
+ grad: its gradient, and the direction the two-loop recursion turns.
65
+ cost: the objective at `x`, which the sufficient decrease is measured from.
66
+ f: forward only, `f(design) -> cost`, so every trial costs one forward
67
+ eval.
68
+ project: applied to every trial, a constraint on the design variables.
69
+
70
+ Returns:
71
+ (design, cost, alpha, trials) of the accepted trial, which is the last one
72
+ attempted even when `alpha` bottomed out at `min_alpha`.
73
+ """
74
+ step = self.step(x, grad) - x
75
+ if not self.pairs:
76
+ step = step * (self.first_step / float(abs(step).max()))
77
+ slope = float(grad.ravel() @ step.ravel())
78
+
79
+ alpha, trials = 1.0, 0
80
+ while True:
81
+ trial = project(x + alpha * step)
82
+ trial_cost = f(trial)
83
+ trials += 1
84
+ if (
85
+ trial_cost <= cost + self.armijo * alpha * slope
86
+ or alpha <= self.min_alpha
87
+ ):
88
+ return trial, trial_cost, alpha, trials
89
+ alpha *= self.shrink
90
+
91
+ def step(self, x: cpt.NDArray, grad: cpt.NDArray) -> cpt.NDArray:
92
+ """Optimization step (evaluated in the line `search`)."""
93
+ flat_x, flat_grad = x.ravel(), grad.ravel()
94
+ if self.x is not None:
95
+ self.put(flat_x - self.x, flat_grad - self.grad)
96
+ self.x, self.grad = flat_x, flat_grad
97
+ return x - self.lr * self.iterate(flat_grad).reshape(x.shape)
98
+
99
+ def put(self, s: cpt.NDArray, y: cpt.NDArray) -> None:
100
+ """Store secant pair `(s, y)` if it satisfies the curvature condition."""
101
+ if y @ s > 0.0: # skip pairs violating the curvature condition
102
+ self.pairs.append((s, y))
103
+
104
+ def iterate(self, q: cpt.NDArray) -> cpt.NDArray:
105
+ """Two-loop recursion: turn gradient `q` into the L-BFGS descent direction."""
106
+ # backward pass over stored pairs (newest to oldest)
107
+ alpha = []
108
+ for s, y in reversed(self.pairs):
109
+ a = (s @ q) / (y @ s)
110
+ q = q - a * y
111
+ alpha.append(a)
112
+
113
+ if self.pairs:
114
+ s, y = self.pairs[-1]
115
+ q = q * ((s @ y) / (y @ y))
116
+
117
+ # forward pass building the descent direction (oldest to newest)
118
+ r = q
119
+ for (s, y), a in zip(self.pairs, reversed(alpha)):
120
+ beta = (y @ r) / (y @ s)
121
+ r = r + (a - beta) * s
122
+
123
+ return r