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)
| 190 | |
| 191 | @torch.compile(dynamic=True, fullgraph=True) |
| 192 | def 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 | |
| 236 | def apply_normuon_variance_reduction(v_chunk, second_momentum_buffer, beta2, red_dim): |
nothing calls this directly
no outgoing calls
no test coverage detected