poincare-hyper / src /physics_losses.py
DHDRL's picture
Rename physics_losses.py to src/physics_losses.py
352610b verified
Raw
History Blame Contribute Delete
3.02 kB
"""
Physics-inspired auxiliary losses for spatiotemporal fields (SciML / PINN style).
These provide soft inductive biases appropriate for reaction-diffusion,
active-matter, and fluid-like systems common in The Well.
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
def spatial_smoothness(fields: torch.Tensor) -> torch.Tensor:
if fields.dim() == 5:
fields = fields.reshape(-1, *fields.shape[2:])
dx = fields[:, :, :, 1:] - fields[:, :, :, :-1]
dy = fields[:, :, 1:, :] - fields[:, :, :-1, :]
return dx.pow(2).mean() + dy.pow(2).mean()
def temporal_consistency(fields: torch.Tensor) -> torch.Tensor:
if fields.dim() != 5 or fields.size(1) < 2:
return fields.new_zeros(())
return (fields[:, 1:] - fields[:, :-1]).pow(2).mean()
def simple_conservation_proxy(fields: torch.Tensor) -> torch.Tensor:
"""
Discrete mass-like conservation per channel.
fields: (B, T, C, H, W)
H_c(t) = sum_{h,w} fields[b, t, c, h, w] (spatial integral)
Loss = mean over (b, c) of Var_t(H_c)
This is the correct zero-parameter conservation residual for
continuity / reaction-diffusion style systems. It constrains the
integrated quantity, not merely the average pixel value.
"""
if fields.dim() != 5 or fields.size(1) < 2:
return fields.new_zeros(())
# Spatial integral per channel: (B, T, C)
H = fields.sum(dim=(-2, -1))
return H.var(dim=1).mean()
def residual_dynamics_penalty(fields: torch.Tensor) -> torch.Tensor:
if fields.dim() != 5 or fields.size(1) < 2:
return fields.new_zeros(())
dt = fields[:, 1:] - fields[:, :-1] # (B, T-1, C, H, W)
dx = dt[:, :, :, :, 1:] - dt[:, :, :, :, :-1]
dy = dt[:, :, :, 1:, :] - dt[:, :, :, :-1, :]
return dx.pow(2).mean() + dy.pow(2).mean()
def channel_coupling_penalty(fields: torch.Tensor) -> torch.Tensor:
if fields.dim() != 5 or fields.size(2) < 2:
return fields.new_zeros(())
dt = fields[:, 1:] - fields[:, :-1]
C = dt.size(2)
channels = [dt[:, :, c].reshape(dt.size(0), -1) for c in range(C)]
channels = [c - c.mean(dim=1, keepdim=True) for c in channels]
stds = [c.std(dim=1) + 1e-6 for c in channels]
corrs = []
for i in range(C):
for j in range(i + 1, C):
corr_ij = (channels[i] * channels[j]).mean(dim=1) / (stds[i] * stds[j])
corrs.append(corr_ij)
mean_abs_corr = torch.stack(corrs, dim=0).abs().mean(dim=0)
return (1.0 - mean_abs_corr).mean()
def combined_physics_loss(
fields: torch.Tensor,
w_smooth: float = 0.01,
w_temp: float = 0.01,
w_cons: float = 0.005,
w_resid: float = 0.005,
w_couple: float = 0.002,
) -> torch.Tensor:
return (
w_smooth * spatial_smoothness(fields)
+ w_temp * temporal_consistency(fields)
+ w_cons * simple_conservation_proxy(fields)
+ w_resid * residual_dynamics_penalty(fields)
+ w_couple * channel_coupling_penalty(fields)
)