MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / polar_express_fn

Function polar_express_fn

optimizers/generic_optim.py:192–233  ·  view source on GitHub ↗

Polar Express Sign Method: https://arxiv.org/pdf/2505.16932 by Noah Amsel, David Persson, Christopher Musco, Robert M. Gower.

(G: torch.Tensor, split_baddbmm: bool = False)

Source from the content-addressed store, hash-verified

190
191@torch.compile(dynamic=True, fullgraph=True)
192def polar_express_fn(G: torch.Tensor, split_baddbmm: bool = False):
193 """
194 Polar Express Sign Method: https://arxiv.org/pdf/2505.16932
195 by Noah Amsel, David Persson, Christopher Musco, Robert M. Gower.
196 """
197 X = G.bfloat16()
198 if G.size(-2) > G.size(-1):
199 X = X.mT
200
201 # Ensure spectral norm is at most 1
202 X = X / (X.norm(dim=(-2, -1), keepdim=True) * (1 + 2e-2) + 1e-6)
203
204 # Allocate buffers
205 X = X.contiguous()
206 C = torch.empty_like(X)
207
208 # Select batched vs unbatched
209 if split_baddbmm:
210 BX_matmul = torch.bmm if X.ndim > 2 else torch.mm
211 else:
212 aX_plus_BX = torch.baddbmm if X.ndim > 2 else torch.addmm
213
214 # Perform the iterations
215 for a, b, c in polar_express_coeffs:
216 A = X @ X.mT
217 B = b * A + c * A @ A
218
219 # Referencing X twice causes pytorch to make a defensive copy,
220 # resulting in a cudaMemcpyAsync in baddbmm.
221 # For large matrices (i.e., the mlp weights), it's faster to split
222 # the operation into two kernels to avoid this.
223 if split_baddbmm:
224 BX_matmul(B, X, out=C) # C = B @ X
225 C.add_(X, alpha=a) # C = C + a*X (in-place, X only read)
226 else:
227 aX_plus_BX(X, B, X, beta=a, out=C) # C = a * X + B @ X
228
229 X, C = C, X # Swap references to avoid unnecessary copies
230
231 if G.size(-2) > G.size(-1):
232 X = X.mT
233 return X
234
235
236def apply_normuon_variance_reduction(v_chunk, second_momentum_buffer, beta2, red_dim):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected