IMF / trace_imf.py
fushinguyenex's picture
Upload completed Trace-iMF training run
cfaf85b verified
Raw History Blame Contribute Delete
5.17 kB
"""Two-trace quadratic objective from idea.pdf, sections 10 and 14."""
import torch
import torch.nn.functional as F
class TraceIMF:
def __init__(self, channels=3, image_size=32, num_classes=10, lambda_diag=1.0):
if lambda_diag <= 0:
raise ValueError("lambda_diag must be positive to train both traces")
self.channels = channels
self.image_size = image_size
self.num_classes = num_classes
self.lambda_diag = lambda_diag
def loss(self, model, x, labels=None, *, derivative_model=None, t=None, noise=None,
r=None, generator=None):
"""x is raw image data in [0,1]. Optional t/noise permit exact checks.
Use the unwrapped model for JVP under Accelerate/DDP. The single joint
grad-enabled call through model trains both edge and diagonal safely.
"""
b = x.shape[0]
x = x.float() * 2 - 1
t = torch.rand(b, device=x.device, generator=generator) if t is None else t.to(x.device).float()
noise = torch.randn(x.shape, device=x.device, generator=generator) if noise is None else noise.to(x.device).float()
if t.shape != (b,) or noise.shape != x.shape:
raise ValueError("Expected t with shape (B,) and noise with the image shape")
if self.num_classes is not None and labels is None:
raise ValueError("Class-conditioned training requires labels")
if self.num_classes is None:
labels = None
t_image = t[:, None, None, None]
z = (1 - t_image) * x + t_image * noise
target = noise - x
r = torch.zeros_like(t) if r is None else r.to(x.device).float()
if r.shape != t.shape or (r < 0).any() or (t > 1).any() or (r > t).any():
raise ValueError("Require 0 <= r <= t <= 1, with r and t shaped (B,)")
if getattr(self, "channels_last", False):
z = z.contiguous(memory_format=torch.channels_last)
# Both branches use the SAME u head, evaluated at r=0 and r=t.
pair_labels = None if labels is None else torch.cat((labels, labels))
predictions = model(torch.cat((z, z)), torch.cat((t, t)),
torch.cat((r, t)), y=pair_labels,
use_flash_attention=True)
edge, diagonal = predictions.chunk(2)
jvp_model = model if derivative_model is None else derivative_model
# FP32 JVP: native attention supports forward AD; SDPA is used only in
# the grad-enabled forward. No parameter gradient through the JVP.
use_bf16 = getattr(self, "jvp_precision", "fp32") == "bf16" and x.device.type == "cuda"
with torch.no_grad(), torch.autocast(device_type=x.device.type, enabled=use_bf16,
dtype=torch.bfloat16):
def edge_fn(z_value, t_value, r_value):
# Accelerate may autocast the instance's forward even after
# unwrap_model. The class forward bypasses that AMP wrapper
# without mutating the grad-enabled training forward.
return type(jvp_model).forward(jvp_model, z_value, t_value, r_value,
y=labels, use_flash_attention=False)
_, derivative = torch.func.jvp(
edge_fn, (z, t, r),
(diagonal.detach().float(), torch.ones_like(t), torch.zeros_like(r)),
)
compound = edge.float() + (t - r)[:, None, None, None] * derivative.detach().float()
edge_loss = F.mse_loss(compound, target)
diag_loss = F.mse_loss(diagonal.float(), target)
loss = edge_loss + self.lambda_diag * diag_loss
return loss, {"edge_loss": edge_loss.detach(), "diag_loss": diag_loss.detach(),
"jvp_rms": derivative.detach().square().mean().sqrt()}
@torch.no_grad()
def sample(self, model, n_samples=None, labels=None, *, device=None, noise=None):
"""Exactly one evaluation at (t,r)=(1,0), then map back to [0,1]."""
was_training = model.training
model.eval()
try:
if device is None:
device = next(model.parameters()).device
if labels is not None:
labels = labels.to(device)
n_samples = len(labels)
if self.num_classes is not None and labels is None:
raise ValueError("Class-conditioned sampling requires labels")
if self.num_classes is None:
labels = None
if noise is None:
if n_samples is None or n_samples <= 0:
raise ValueError("n_samples must be positive")
noise = torch.randn(n_samples, self.channels, self.image_size,
self.image_size, device=device)
else:
noise = noise.to(device)
t = torch.ones(noise.shape[0], device=device)
u = model(noise, t, torch.zeros_like(t), y=labels,
use_flash_attention=True)
return ((noise - u).clamp(-1, 1) + 1) * 0.5
finally:
model.train(was_training)