compdiff-chest-xray / compdiff2.py
mahmoudibra98's picture
Release CompDiff chest v2: typed 3-attribute conditioner (age in HCN), step-10k checkpoint
af2d273 verified
Raw History Blame Contribute Delete
26.4 kB
"""
CompDiff conditioner: Typed Compositional Conditioner (released standalone module).
This file is a verbatim copy of the conditioner class used to train the released
checkpoint (`roentgenv2/train_code/compdiff2.py` in the CompDiff repository), with
the training-only pieces (loss, config-driven builder, self-tests) removed and the
`save_pretrained` / `from_pretrained` helpers switched to safetensors.
Design:
* Typed encoders: sex / race = nn.Embedding (nominal);
age = continuous years -> sinusoidal features -> MLP (ordinal).
* Composer: pairwise-MLP hierarchy (age x sex, age x race, sex x race -> all),
then each attribute is re-contextualised against the composed state.
* Output: 4 tokens (t_age, t_sex, t_race, t_cls) in the UNet cross-attention
space (d_ctx = 1024 for SD 2.1), concatenated to the 77 CLIP text tokens.
* Aux heads on the output tokens (sex CE, race CE, age regression, joint CE)
were used during training; their weights are kept so the module loads the
checkpoint strictly, but they are not needed for generation.
Interface:
forward(sex_idx [B], race_idx [B], age_continuous [B] float years)
-> (ctx [B, T, d_ctx], mu [B, d_node], logsigma [B, d_node],
aux_logits dict | None, time_emb [B, d_time_emb] | None)
Eval mode is deterministic (z = mu).
"""
import math
import json
import os
from typing import Tuple, Optional, Dict
import torch
import torch.nn as nn
import torch.nn.functional as F
class MLP(nn.Module):
"""LayerNorm -> Linear -> SiLU -> Dropout -> Linear (same block as CompDiff-1)."""
def __init__(self, d_in: int, d_hidden: int, d_out: int, dropout: float = 0.1):
super().__init__()
self.net = nn.Sequential(
nn.LayerNorm(d_in),
nn.Linear(d_in, d_hidden),
nn.SiLU(),
nn.Dropout(dropout),
nn.Linear(d_hidden, d_out),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
def sinusoidal_age_features(age_years: torch.Tensor, dim: int, max_period: float = 10000.0) -> torch.Tensor:
"""
Sinusoidal features of age in years — the same encoding family the UNet
uses for the diffusion timestep, giving age smooth ordinal geometry by
construction (nearby ages -> nearby features).
Args:
age_years: [B] float tensor of ages in years
dim: feature dimension (must be even)
Returns:
[B, dim] float tensor
"""
half = dim // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(half, dtype=torch.float32, device=age_years.device)
/ half
)
args = age_years.float().unsqueeze(-1) * freqs.unsqueeze(0) # [B, half]
return torch.cat([torch.cos(args), torch.sin(args)], dim=-1) # [B, dim]
class CompDiff2Conditioner(nn.Module):
"""
Typed compositional demographic conditioner (CompDiff-2).
Args:
num_sex, num_race: category counts for the nominal attributes
num_age_bins: bin count kept ONLY for the joint-cell aux head and
monitoring (age itself is continuous inside the conditioner)
d_node: composer latent dimension
d_ctx: UNet cross-attention dimension (1024 for SD 2.1)
d_time_emb: UNet timestep-embedding dimension (1280 for SD 2.1)
max_age: normalization constant for the age regression target
age_freq_dim: sinusoidal feature dimension for the age encoder
composer: 'hierarchical' | 'transformer'
multi_token: single fused token (False) vs per-attribute tokens (True)
route_b: emit a zero-init timestep-embedding modulation vector
num_registers: extra unsupervised register tokens (multi_token only)
attr_dropout_prob: per-sample per-attribute prob of replacing an
attribute with its learned null embedding (training only)
full_dropout_prob: per-sample prob of dropping ALL attributes at once
(training only; CFG-style unconditional demographic branch)
use_uncertainty: variational latent on the composed representation
use_aux_loss: build aux heads on the output tokens
aux_hidden_dim: hidden dim of aux heads
dropout: dropout inside MLPs / transformer
transformer_layers, transformer_heads: composer size ('transformer')
flat_hidden: hidden width of the 'flat' composer blocks. 664 matches the
hierarchical multi-token composer's parameter count within 0.1%
(2,898,632 vs 2,896,640) at d_node=256 -- see composer_num_params().
"""
def __init__(
self,
num_sex: int = 2,
num_race: int = 4,
num_age_bins: int = 5,
d_node: int = 256,
d_ctx: int = 1024,
d_time_emb: int = 1280,
max_age: float = 100.0,
age_freq_dim: int = 128,
composer: str = "hierarchical",
multi_token: bool = False,
route_b: bool = False,
num_registers: int = 0,
attr_dropout_prob: float = 0.0,
full_dropout_prob: float = 0.0,
use_uncertainty: bool = True,
use_aux_loss: bool = True,
aux_hidden_dim: int = 512,
dropout: float = 0.1,
transformer_layers: int = 2,
transformer_heads: int = 4,
flat_hidden: int = 664,
):
super().__init__()
assert composer in ("hierarchical", "transformer", "flat"), f"Unknown composer: {composer}"
assert age_freq_dim % 2 == 0, "age_freq_dim must be even"
self.config = {
"num_sex": num_sex,
"num_race": num_race,
"num_age_bins": num_age_bins,
"d_node": d_node,
"d_ctx": d_ctx,
"d_time_emb": d_time_emb,
"max_age": max_age,
"age_freq_dim": age_freq_dim,
"composer": composer,
"multi_token": multi_token,
"route_b": route_b,
"num_registers": num_registers,
"attr_dropout_prob": attr_dropout_prob,
"full_dropout_prob": full_dropout_prob,
"use_uncertainty": use_uncertainty,
"use_aux_loss": use_aux_loss,
"aux_hidden_dim": aux_hidden_dim,
"dropout": dropout,
"transformer_layers": transformer_layers,
"transformer_heads": transformer_heads,
"flat_hidden": flat_hidden,
}
self.num_sex = num_sex
self.num_race = num_race
self.num_age_bins = num_age_bins
self.d_node = d_node
self.d_ctx = d_ctx
self.d_time_emb = d_time_emb
self.max_age = float(max_age)
self.age_freq_dim = age_freq_dim
self.composer_type = composer
self.multi_token = multi_token
self.route_b = route_b
self.num_registers = num_registers if multi_token else 0
self.attr_dropout_prob = attr_dropout_prob
self.full_dropout_prob = full_dropout_prob
self.use_uncertainty = use_uncertainty
self.use_aux_loss = use_aux_loss
# Age is always inside the composer for CompDiff-2 (that is the point);
# kept as an attribute for pipeline code that introspects it.
self.encode_age = True
# === Typed attribute encoders ===
self.emb_sex = nn.Embedding(num_sex, d_node)
self.emb_race = nn.Embedding(num_race, d_node)
self.age_encoder = nn.Sequential(
nn.Linear(age_freq_dim, d_node),
nn.SiLU(),
nn.Linear(d_node, d_node),
)
# Learned null embeddings ("attribute unspecified") for dropout and
# partial conditioning at inference.
self.null_age = nn.Parameter(torch.zeros(d_node))
self.null_sex = nn.Parameter(torch.zeros(d_node))
self.null_race = nn.Parameter(torch.zeros(d_node))
# === Composer ===
if composer == "hierarchical":
# CompDiff-1 topology with typed inputs (stage 2a-2c)
self.compose_age_sex = MLP(2 * d_node, 2 * d_node, d_node, dropout)
self.compose_age_race = MLP(2 * d_node, 2 * d_node, d_node, dropout)
self.compose_sex_race = MLP(2 * d_node, 2 * d_node, d_node, dropout)
self.compose_all = MLP(3 * d_node, 2 * d_node, d_node, dropout)
if multi_token:
# Contextualize each attribute against the composed child so
# attribute tokens are "attribute-in-context" representations.
self.ctx_age = MLP(2 * d_node, 2 * d_node, d_node, dropout)
self.ctx_sex = MLP(2 * d_node, 2 * d_node, d_node, dropout)
self.ctx_race = MLP(2 * d_node, 2 * d_node, d_node, dropout)
elif composer == "flat":
# Parameter/depth-matched NON-compositional control (review item 3).
# Same three-stage MLP pipeline as 'hierarchical' (pair-level ->
# compose_all -> per-attribute contextualization), same block type
# (LN -> Linear -> SiLU -> Dropout -> Linear), same depth (6 linear
# layers to the attribute tokens, 4 to h_demo), but every stage is a
# single MLP over the FULL concatenation: no pairwise factorization,
# no per-attribute routing. Widths are tuned (flat_hidden) so the
# composer parameter count matches 'hierarchical' within ~0.1%.
# stage 1: [e_age, e_sex, e_race] (3d) -> H -> 3d (~ 3 pair MLPs)
# stage 2: 3d -> H -> d = h_demo (~ compose_all)
# stage 3: [e_age, e_sex, e_race, h_demo] (4d) -> H -> 3d,
# split into (c_age, c_sex, c_race) (~ ctx_age/sex/race)
H = int(flat_hidden)
self.flat_stage1 = MLP(3 * d_node, H, 3 * d_node, dropout)
self.flat_stage2 = MLP(3 * d_node, H, d_node, dropout)
if multi_token:
self.flat_stage3 = MLP(4 * d_node, H, 3 * d_node, dropout)
else:
# Transformer composer (stage 2d): [t_age, t_sex, t_race, CLS, regs]
self.cls_token = nn.Parameter(torch.zeros(d_node))
num_slots = 4 + self.num_registers
self.type_emb = nn.Parameter(torch.zeros(num_slots, d_node))
if self.num_registers > 0:
self.register_tokens = nn.Parameter(torch.zeros(self.num_registers, d_node))
enc_layer = nn.TransformerEncoderLayer(
d_model=d_node,
nhead=transformer_heads,
dim_feedforward=2 * d_node,
dropout=dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.composer = nn.TransformerEncoder(enc_layer, num_layers=transformer_layers)
# === Variational latent on the composed representation ===
if use_uncertainty:
self.mu_head = nn.Linear(d_node, d_node)
self.logsigma_head = nn.Linear(d_node, d_node)
# === Projections to cross-attention space ===
def make_proj():
return nn.Sequential(nn.LayerNorm(d_node), nn.Linear(d_node, d_ctx))
self.proj_cls = make_proj()
if multi_token:
self.proj_age = make_proj()
self.proj_sex = make_proj()
self.proj_race = make_proj()
if self.num_registers > 0:
self.proj_reg = make_proj()
# === Route B: zero-init projection into the timestep embedding ===
if route_b:
self.time_proj = nn.Linear(d_node, d_time_emb)
nn.init.zeros_(self.time_proj.weight)
nn.init.zeros_(self.time_proj.bias)
# === Aux heads ON OUTPUT TOKENS (post-projection, d_ctx) ===
if use_aux_loss:
def make_head(d_out):
return nn.Sequential(
nn.LayerNorm(d_ctx),
nn.Linear(d_ctx, aux_hidden_dim),
nn.SiLU(),
nn.Dropout(dropout),
nn.Linear(aux_hidden_dim, d_out),
)
self.sex_classifier = make_head(num_sex)
self.race_classifier = make_head(num_race)
self.age_regressor = make_head(1)
self.joint_classifier = make_head(num_age_bins * num_sex * num_race)
self._init_weights()
# ------------------------------------------------------------------
@property
def num_output_tokens(self) -> int:
if not self.multi_token:
return 1
return 4 + self.num_registers # age, sex, race, cls (+ registers)
def composer_num_params(self) -> int:
"""Parameter count of the COMPOSER ONLY (everything between the typed
attribute embeddings and the variational/projection heads). Used to
parameter-match the 'flat' control against 'hierarchical'."""
if self.composer_type == "hierarchical":
mods = [self.compose_age_sex, self.compose_age_race, self.compose_sex_race, self.compose_all]
if self.multi_token:
mods += [self.ctx_age, self.ctx_sex, self.ctx_race]
elif self.composer_type == "flat":
mods = [self.flat_stage1, self.flat_stage2]
if self.multi_token:
mods.append(self.flat_stage3)
else:
mods = [self.composer]
extra = self.cls_token.numel() + self.type_emb.numel()
if self.num_registers > 0:
extra += self.register_tokens.numel()
return sum(p.numel() for m in mods for p in m.parameters()) + extra
return sum(p.numel() for m in mods for p in m.parameters())
def composer_depth(self) -> int:
"""Number of nn.Linear layers on the longest input->output path of the composer."""
if self.composer_type == "hierarchical":
return 6 if self.multi_token else 4
if self.composer_type == "flat":
return 6 if self.multi_token else 4
return 4 * self.config["transformer_layers"] # per layer: attn in_proj, out_proj, FFN x2
def _init_weights(self):
for emb in (self.emb_sex, self.emb_race):
nn.init.normal_(emb.weight, mean=0.0, std=0.02)
for p in (self.null_age, self.null_sex, self.null_race):
nn.init.normal_(p, mean=0.0, std=0.02)
if self.composer_type == "transformer":
nn.init.normal_(self.cls_token, mean=0.0, std=0.02)
nn.init.normal_(self.type_emb, mean=0.0, std=0.02)
if self.num_registers > 0:
nn.init.normal_(self.register_tokens, mean=0.0, std=0.02)
if self.use_uncertainty:
nn.init.normal_(self.mu_head.weight, mean=0.0, std=0.01)
nn.init.zeros_(self.mu_head.bias)
nn.init.normal_(self.logsigma_head.weight, mean=0.0, std=0.01)
nn.init.constant_(self.logsigma_head.bias, -1.0)
# ------------------------------------------------------------------
def _encode_attributes(
self,
sex_idx: torch.Tensor,
race_idx: torch.Tensor,
age_continuous: Optional[torch.Tensor],
apply_dropout: bool,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Typed grandparent embeddings, with optional null-dropout."""
e_sex = self.emb_sex(sex_idx)
e_race = self.emb_race(race_idx)
B = e_sex.shape[0]
if age_continuous is not None:
feats = sinusoidal_age_features(age_continuous, self.age_freq_dim)
e_age = self.age_encoder(feats.to(dtype=e_sex.dtype))
else:
# Partial conditioning: age unspecified
e_age = self.null_age.unsqueeze(0).expand(B, -1).to(dtype=e_sex.dtype)
if apply_dropout and self.training and (self.attr_dropout_prob > 0 or self.full_dropout_prob > 0):
device = e_sex.device
full = torch.rand(B, device=device) < self.full_dropout_prob
for name, null in (("age", self.null_age), ("sex", self.null_sex), ("race", self.null_race)):
drop = (torch.rand(B, device=device) < self.attr_dropout_prob) | full
mask = drop.unsqueeze(-1).to(dtype=e_sex.dtype)
null_row = null.unsqueeze(0).to(dtype=e_sex.dtype)
if name == "age":
e_age = e_age * (1 - mask) + null_row * mask
elif name == "sex":
e_sex = e_sex * (1 - mask) + null_row * mask
else:
e_race = e_race * (1 - mask) + null_row * mask
return e_age, e_sex, e_race
@torch.no_grad()
def forward_unconditional(self, batch_size: int, device=None, dtype=None):
"""Tokens (+ Route B vector) for the model's TRAINED 'demographics
unspecified' state: all three attributes at their learned null
embeddings.
Only meaningful for models trained with attribute dropout (stage 2e);
for the others the null embeddings never received gradient and this is
not a trained state. Deterministic (z = mu, no sampling), so it is safe
as the unconditional branch of classifier-free guidance.
Add-only: does not touch forward() semantics.
"""
device = device or self.null_sex.device
dtype = dtype or self.null_sex.dtype
exp = lambda p: p.unsqueeze(0).expand(batch_size, -1).to(device=device, dtype=dtype)
e_age, e_sex, e_race = exp(self.null_age), exp(self.null_sex), exp(self.null_race)
h_demo, attr_ctx = self._compose(e_age, e_sex, e_race)
z = self.mu_head(h_demo) if self.use_uncertainty else h_demo
t_cls = self.proj_cls(z)
if self.multi_token:
tokens = [
self.proj_age(attr_ctx["age"]),
self.proj_sex(attr_ctx["sex"]),
self.proj_race(attr_ctx["race"]),
t_cls,
]
if self.num_registers > 0:
if self.composer_type == "transformer":
regs = attr_ctx["registers"]
else:
regs = self.register_tokens.unsqueeze(0).expand(batch_size, -1, -1)
tokens.extend(self.proj_reg(regs[:, r]) for r in range(self.num_registers))
ctx = torch.stack(tokens, dim=1)
else:
ctx = t_cls.unsqueeze(1)
time_emb = self.time_proj(z) if self.route_b else None
return ctx, time_emb
def _compose(
self,
e_age: torch.Tensor,
e_sex: torch.Tensor,
e_race: torch.Tensor,
) -> Tuple[torch.Tensor, Optional[Dict[str, torch.Tensor]]]:
"""
Run the composer.
Returns:
h_demo: [B, d_node] composed representation
attr_ctx: dict of contextualized per-attribute states [B, d_node]
(None when multi_token=False)
"""
if self.composer_type == "hierarchical":
h_age_sex = self.compose_age_sex(torch.cat([e_age, e_sex], dim=-1))
h_age_race = self.compose_age_race(torch.cat([e_age, e_race], dim=-1))
h_sex_race = self.compose_sex_race(torch.cat([e_sex, e_race], dim=-1))
h_demo = self.compose_all(torch.cat([h_age_sex, h_age_race, h_sex_race], dim=-1))
attr_ctx = None
if self.multi_token:
attr_ctx = {
"age": self.ctx_age(torch.cat([e_age, h_demo], dim=-1)),
"sex": self.ctx_sex(torch.cat([e_sex, h_demo], dim=-1)),
"race": self.ctx_race(torch.cat([e_race, h_demo], dim=-1)),
}
return h_demo, attr_ctx
elif self.composer_type == "flat":
x = torch.cat([e_age, e_sex, e_race], dim=-1)
h1 = self.flat_stage1(x)
h_demo = self.flat_stage2(h1)
attr_ctx = None
if self.multi_token:
c = self.flat_stage3(torch.cat([x, h_demo], dim=-1))
c_age, c_sex, c_race = torch.split(c, self.d_node, dim=-1)
attr_ctx = {"age": c_age, "sex": c_sex, "race": c_race}
return h_demo, attr_ctx
else:
B = e_age.shape[0]
seq = [e_age, e_sex, e_race, self.cls_token.unsqueeze(0).expand(B, -1)]
if self.num_registers > 0:
for r in range(self.num_registers):
seq.append(self.register_tokens[r].unsqueeze(0).expand(B, -1))
x = torch.stack(seq, dim=1) # [B, S, d_node]
x = x + self.type_emb.unsqueeze(0)
out = self.composer(x)
h_demo = out[:, 3] # CLS position
attr_ctx = None
if self.multi_token:
attr_ctx = {"age": out[:, 0], "sex": out[:, 1], "race": out[:, 2]}
if self.num_registers > 0:
attr_ctx["registers"] = out[:, 4:]
return h_demo, attr_ctx
# ------------------------------------------------------------------
def forward(
self,
sex_idx: torch.Tensor,
race_idx: torch.Tensor,
age_continuous: Optional[torch.Tensor] = None,
age_idx: Optional[torch.Tensor] = None, # accepted for interface compat; unused
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, Optional[Dict[str, torch.Tensor]], Optional[torch.Tensor]]:
e_age, e_sex, e_race = self._encode_attributes(
sex_idx, race_idx, age_continuous, apply_dropout=True
)
h_demo, attr_ctx = self._compose(e_age, e_sex, e_race)
# DDP: the learned null embeddings are only consumed on dropout /
# partial-conditioning batches; tie them into the graph with a
# zero-scaled anchor so every parameter produces a gradient on every
# step (otherwise DDP's reducer errors with "parameters that were not
# used in producing loss" — observed as SLURM job 122136, indices 0-2).
h_demo = h_demo + 0.0 * (self.null_age + self.null_sex + self.null_race).sum()
# Variational latent on the composed (CLS) representation only
if self.use_uncertainty:
mu = self.mu_head(h_demo)
logsigma = torch.clamp(self.logsigma_head(h_demo), min=-5.0, max=1.0)
if self.training:
z = mu + torch.exp(logsigma) * torch.randn_like(mu)
else:
z = mu
else:
mu = h_demo
logsigma = torch.zeros_like(h_demo)
z = h_demo
# Output tokens for cross-attention (Route A)
t_cls = self.proj_cls(z) # [B, d_ctx]
if self.multi_token:
tokens = [
self.proj_age(attr_ctx["age"]),
self.proj_sex(attr_ctx["sex"]),
self.proj_race(attr_ctx["race"]),
t_cls,
]
if self.num_registers > 0:
if self.composer_type == "transformer":
regs = attr_ctx["registers"] # [B, R, d_node]
else:
B = t_cls.shape[0]
regs = self.register_tokens.unsqueeze(0).expand(B, -1, -1)
tokens.extend([self.proj_reg(regs[:, r]) for r in range(self.num_registers)])
ctx = torch.stack(tokens, dim=1) # [B, T, d_ctx]
else:
ctx = t_cls.unsqueeze(1) # [B, 1, d_ctx]
# Route B: global modulation vector for the timestep embedding
time_emb = self.time_proj(z) if self.route_b else None
# Aux logits from the tokens the UNet actually sees
aux_logits = None
if self.use_aux_loss:
if self.multi_token:
tok_age, tok_sex, tok_race = ctx[:, 0], ctx[:, 1], ctx[:, 2]
tok_joint = ctx[:, 3]
else:
tok_age = tok_sex = tok_race = tok_joint = ctx[:, 0]
aux_logits = {
"sex": self.sex_classifier(tok_sex),
"race": self.race_classifier(tok_race),
"age_pred": self.age_regressor(tok_age).squeeze(-1), # normalized age
"joint": self.joint_classifier(tok_joint),
}
return ctx, mu, logsigma, aux_logits, time_emb
# ------------------------------------------------------------------
def compute_compositional_loss(
self,
sex_idx: torch.Tensor,
race_idx: torch.Tensor,
age_continuous: Optional[torch.Tensor] = None,
age_idx: Optional[torch.Tensor] = None, # interface compat; unused
) -> torch.Tensor:
"""Soft additive anchor: cos(h_demo, e_age + e_sex + e_race)."""
e_age, e_sex, e_race = self._encode_attributes(
sex_idx, race_idx, age_continuous, apply_dropout=False
)
h_demo, _ = self._compose(e_age, e_sex, e_race)
h_additive = e_age + e_sex + e_race
cos_sim = F.cosine_similarity(h_demo, h_additive, dim=-1)
return (1 - cos_sim).mean()
def get_uncertainty(
self,
sex_idx: torch.Tensor,
race_idx: torch.Tensor,
age_continuous: Optional[torch.Tensor] = None,
) -> torch.Tensor:
_, _, logsigma, _, _ = self.forward(sex_idx, race_idx, age_continuous=age_continuous)
return torch.exp(logsigma).mean(dim=-1)
# ------------------------------------------------------------------
def save_pretrained(self, save_dir: str):
os.makedirs(save_dir, exist_ok=True)
with open(os.path.join(save_dir, "config.json"), "w") as f:
json.dump(self.config, f, indent=2)
from safetensors.torch import save_file
save_file({k: v.contiguous() for k, v in self.state_dict().items()},
os.path.join(save_dir, "model.safetensors"), metadata={"format": "pt"})
@classmethod
def from_pretrained(cls, save_dir: str, device: str = "cpu"):
with open(os.path.join(save_dir, "config.json"), "r") as f:
config = json.load(f)
model = cls(**config)
st_path = os.path.join(save_dir, "model.safetensors")
if os.path.exists(st_path):
from safetensors.torch import load_file
state_dict = load_file(st_path, device="cpu")
else: # legacy pickle layout
state_dict = torch.load(os.path.join(save_dir, "pytorch_model.bin"), map_location="cpu")
model.load_state_dict(state_dict, strict=True)
model.to(device)
model.eval()
return model