Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,48 @@ cosmodiff_sample.py --config path/to/config.yaml \
--num_steps 25
```

## Fresh DiT zero initialization

To opt a new DiT training run into zero modulation/output initialization, add
`initialization: adaln_zero` alongside `class` and `kwargs` in the model config:

```yaml
model:
class: DiTTransformer2DModel
initialization: adaln_zero
kwargs:
sample_size: 128
patch_size: 8
in_channels: 1
out_channels: 1
num_layers: 16
num_attention_heads: 12
attention_head_dim: 64
num_embeds_ada_norm: 1000
norm_type: ada_norm_zero
```

Keep the architecture arguments appropriate for your run; the policy does not
change them. It zeroes the weight and bias of every block's `norm1.linear`
(attention/FFN shifts, scales, and residual gates), plus `proj_out_1` (final
conditioning) and `proj_out_2` (prediction head). Initially, the blocks are
identity mappings and the denoising prediction is zero; these values are not
frozen and training learns them. Attention, FFN, and embedding weights retain
their native initialization.

This implements the zero-modulation/output part of
[Peebles and Xie's DiT initialization](https://arxiv.org/abs/2212.09748), not
the complete reference initialization or training recipe. Omitted
`initialization`, or `initialization: native`, preserves existing behavior.
Only native `DiTTransformer2DModel` with `norm_type: ada_norm_zero` is supported;
incompatible projection APIs fail before any weights are zeroed. This policy
was checked with diffusers 0.38.0 without imposing a version pin.

Initialization is applied only when constructing a fresh model. Resuming or
sampling a checkpoint does **not** zero its trained weights. Enabling this
option does not retrofit an already-trained checkpoint or guarantee improved
generation for every dataset or model depth.

## Authors

- Nicholas Kern
Expand Down
1 change: 1 addition & 0 deletions cosmodiff/data/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ augmentations:
# --- model ---
model:
class: UNet2DModel
initialization: native # DiT only: 'adaln_zero' opts in to fresh zero modulation/output init
kwargs:
sample_size: 64
in_channels: 1
Expand Down
62 changes: 62 additions & 0 deletions cosmodiff/initialization.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
"""Explicit, fresh-model initialization policies."""


def initialize_dit_adaln_zero(model, *, fresh: bool) -> dict:
"""Zero a fresh native DiT's modulation and final output projections.

This adopts the zero-modulation/output part of the original DiT recipe,
not its complete initialization scheme. Attention, FFN, patch embedding,
and timestep/class embedding parameters retain diffusers' initialization.
All parameters remain trainable. Never apply this to a loaded checkpoint.

The projection API is validated in full before any parameter is changed.
Unsupported architectures fail explicitly rather than receiving a partial
initialization. The marker prevents applying this policy twice in memory.

Returns:
dict: Names and count of the initialized projections, for inspection.
"""
import torch
from diffusers import DiTTransformer2DModel
from diffusers.models.normalization import AdaLayerNormZero

if not fresh or getattr(model, "_cosmodiff_adaln_zero_initialized", False):
raise RuntimeError("adaLN-Zero initialization requires a fresh, uninitialized model")
if type(model) is not DiTTransformer2DModel:
raise TypeError("adaLN-Zero initialization supports only native DiTTransformer2DModel")

width = model.config.num_attention_heads * model.config.attention_head_dim
if len(model.transformer_blocks) != model.config.num_layers:
raise ValueError("DiT layer count differs from its config")

targets = []
for index, block in enumerate(model.transformer_blocks):
if type(block.norm1) is not AdaLayerNormZero:
raise TypeError(f"DiT block {index} requires norm_type='ada_norm_zero'")
targets.append((
f"transformer_blocks.{index}.norm1.linear", block.norm1.linear, 6 * width,
))
targets.extend([
("proj_out_1", model.proj_out_1, 2 * width),
("proj_out_2", model.proj_out_2, model.config.patch_size**2 * model.out_channels),
])

for name, layer, outputs in targets:
if (type(layer) is not torch.nn.Linear or layer.bias is None
or tuple(layer.weight.shape) != (outputs, width)
or tuple(layer.bias.shape) != (outputs,)
or layer.weight.is_meta or layer.bias.is_meta):
raise ValueError(f"Unsupported DiT projection API: {name}")

with torch.no_grad():
for _, layer, _ in targets:
layer.weight.zero_()
layer.bias.zero_()
model._cosmodiff_adaln_zero_initialized = True

return {
"kind": "adaln_zero_modulation_and_output",
"fresh_only": True,
"linear_modules": [name for name, _, _ in targets],
"zeroed_tensors": 2 * len(targets),
}
247 changes: 247 additions & 0 deletions cosmodiff/tests/test_initialization.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,247 @@
"""Regression tests for opt-in, fresh-only DiT zero initialization."""
import copy
import importlib.util
import sys
from pathlib import Path

import pytest
import torch
from diffusers import DiTTransformer2DModel, UNet2DModel

from cosmodiff.initialization import initialize_dit_adaln_zero
from cosmodiff.utils import parse_config_model


def _dit_kwargs(patch_size=4, num_layers=2):
return dict(
sample_size=8, patch_size=patch_size, in_channels=1, out_channels=1,
num_layers=num_layers, num_attention_heads=2, attention_head_dim=8,
num_embeds_ada_norm=10, norm_type="ada_norm_zero",
)


def _state(model):
return {name: value.clone() for name, value in model.state_dict().items()}


@pytest.mark.parametrize("patch_size,num_layers", [(4, 2), (8, 16)])
def test_only_modulation_and_output_parameters_are_zeroed(patch_size, num_layers):
model = DiTTransformer2DModel(**_dit_kwargs(patch_size, num_layers))
before = _state(model)
report = initialize_dit_adaln_zero(model, fresh=True)
targets = {
f"{name}.{parameter}"
for name in report["linear_modules"] for parameter in ("weight", "bias")
}
assert report["zeroed_tensors"] == 2 * (num_layers + 2)
assert len(targets) == report["zeroed_tensors"]
assert all(parameter.requires_grad for parameter in model.parameters())
for name, value in model.state_dict().items():
if name in targets:
assert torch.count_nonzero(value) == 0, name
else:
assert torch.equal(value, before[name]), name


def test_zero_prediction_and_output_head_can_learn():
torch.manual_seed(123)
model = DiTTransformer2DModel(**_dit_kwargs())
initialize_dit_adaln_zero(model, fresh=True)
images = torch.randn(2, 1, 8, 8)
timesteps = torch.tensor([1, 2])
labels = torch.zeros(2, dtype=torch.long)
output = model(images, timesteps, class_labels=labels).sample
assert torch.count_nonzero(output) == 0
optimizer = torch.optim.SGD(model.parameters(), lr=0.05)
loss = (output - torch.ones_like(output)).square().mean()
loss.backward()
assert torch.isfinite(model.proj_out_2.weight.grad).all()
assert torch.count_nonzero(model.proj_out_2.weight.grad) > 0
optimizer.step()
learned = model(images, timesteps, class_labels=labels).sample
assert torch.isfinite(learned).all()
assert torch.count_nonzero(learned) > 0


def test_initialized_blocks_are_identity_mappings():
model = DiTTransformer2DModel(**_dit_kwargs())
initialize_dit_adaln_zero(model, fresh=True)
features = torch.randn(2, 4, 16)
timesteps = torch.tensor([1, 2])
labels = torch.zeros(2, dtype=torch.long)
for block in model.transformer_blocks:
output = block(features, timestep=timesteps, class_labels=labels)
assert torch.equal(output, features)


@pytest.mark.parametrize("fresh", [False, True])
def test_reinitialization_is_rejected_without_mutation(fresh):
model = DiTTransformer2DModel(**_dit_kwargs())
if fresh:
initialize_dit_adaln_zero(model, fresh=True)
before = _state(model)
with pytest.raises(RuntimeError, match="fresh"):
initialize_dit_adaln_zero(model, fresh=fresh)
for name, value in model.state_dict().items():
assert torch.equal(value, before[name])


def test_incompatible_projection_does_not_partially_zero_model():
model = DiTTransformer2DModel(**_dit_kwargs())
model.proj_out_1 = torch.nn.Linear(16, 1)
before = _state(model)
with pytest.raises(ValueError, match="proj_out_1"):
initialize_dit_adaln_zero(model, fresh=True)
assert not getattr(model, "_cosmodiff_adaln_zero_initialized", False)
for name, value in model.state_dict().items():
assert torch.equal(value, before[name])


def test_wrong_normalization_is_rejected():
model = DiTTransformer2DModel(**_dit_kwargs())
model.transformer_blocks[-1].norm1 = torch.nn.LayerNorm(16)
before = _state(model)
with pytest.raises(TypeError, match="ada_norm_zero"):
initialize_dit_adaln_zero(model, fresh=True)
for name, value in model.state_dict().items():
assert torch.equal(value, before[name])


def test_wrong_model_class_is_rejected():
with pytest.raises(TypeError, match="DiTTransformer2DModel"):
initialize_dit_adaln_zero(torch.nn.Linear(2, 2), fresh=True)


@pytest.mark.parametrize("explicit_native", [False, True])
def test_default_dit_initialization_is_unchanged(explicit_native):
config = {"model": {"class": "DiTTransformer2DModel", "kwargs": _dit_kwargs()}}
if explicit_native:
config["model"]["initialization"] = "native"
torch.manual_seed(123)
expected = DiTTransformer2DModel(**_dit_kwargs())
torch.manual_seed(123)
actual = parse_config_model(config)["model"]
for name, value in actual.state_dict().items():
assert torch.equal(value, expected.state_dict()[name])


def test_config_opt_in_initializes_before_optimizer_without_changing_config():
config = {
"model": {"class": "DiTTransformer2DModel", "initialization": "adaln_zero",
"kwargs": _dit_kwargs()},
"optimizer": {"class": "AdamW", "kwargs": {"lr": 1e-4}},
}
original = copy.deepcopy(config)
result = parse_config_model(config)
model = result["model"]
assert model._cosmodiff_adaln_zero_initialized
assert torch.count_nonzero(model.proj_out_2.weight) == 0
assert config == original
parameters = result["optimizer"].param_groups[0]["params"]
assert {id(parameter) for parameter in parameters} == {
id(parameter) for parameter in model.parameters()
}


@pytest.mark.parametrize("initialization", ["unknown", None])
def test_unknown_initialization_policy_is_rejected(initialization):
config = {"model": {"class": "DiTTransformer2DModel", "initialization": initialization,
"kwargs": _dit_kwargs()}}
with pytest.raises(ValueError, match="model.initialization"):
parse_config_model(config)


def test_config_rejects_zero_init_for_unet():
config = {"model": {"class": "UNet2DModel", "initialization": "adaln_zero"}}
with pytest.raises(ValueError, match="requires DiTTransformer2DModel"):
parse_config_model(config)


def test_unet_native_initialization_is_unchanged():
kwargs = dict(sample_size=8, in_channels=1, out_channels=1, layers_per_block=1,
block_out_channels=(8,), down_block_types=("DownBlock2D",),
up_block_types=("UpBlock2D",), norm_num_groups=4)
torch.manual_seed(123)
expected = UNet2DModel(**kwargs)
torch.manual_seed(123)
actual = parse_config_model({"model": {"class": "UNet2DModel", "kwargs": kwargs}})["model"]
for name, value in actual.state_dict().items():
assert torch.equal(value, expected.state_dict()[name])


def test_save_load_preserves_learned_nonzero_projections(tmp_path):
model = parse_config_model({"model": {
"class": "DiTTransformer2DModel", "initialization": "adaln_zero",
"kwargs": _dit_kwargs(),
}})["model"]
with torch.no_grad():
model.proj_out_2.weight.fill_(0.125)
model.proj_out_2.bias.fill_(0.25)
model.transformer_blocks[0].norm1.linear.weight.fill_(0.375)
model.save_pretrained(tmp_path)
loaded = DiTTransformer2DModel.from_pretrained(tmp_path)
for name, value in loaded.state_dict().items():
assert torch.equal(value, model.state_dict()[name])
with pytest.raises(RuntimeError, match="fresh"):
initialize_dit_adaln_zero(loaded, fresh=False)
assert torch.all(loaded.proj_out_2.weight == 0.125)


def test_checkpoint_loader_preserves_learned_nonzero_projections(tmp_path):
import yaml
from diffusers import DDPMScheduler
from cosmodiff.utils import load_checkpoint

model = DiTTransformer2DModel(**_dit_kwargs())
initialize_dit_adaln_zero(model, fresh=True)
with torch.no_grad():
model.proj_out_2.weight.fill_(0.125)
model.transformer_blocks[0].norm1.linear.bias.fill_(0.25)
model.save_pretrained(tmp_path)
DDPMScheduler(num_train_timesteps=10).save_pretrained(tmp_path)
checkpoint_config = {
"model": {"initialization": "adaln_zero"},
"noise_scheduler": {"class": "diffusers.DDPMScheduler"},
"optimizer": {"class": "torch.optim.AdamW"},
"lr_scheduler": {"class": "torch.optim.lr_scheduler.ConstantLR",
"kwargs": {"factor": 1.0, "total_iters": 0}},
}
(tmp_path / "checkpoint_config.yaml").write_text(yaml.safe_dump(checkpoint_config))
loaded, _, optimizer, _, _ = load_checkpoint(str(tmp_path))
for name, value in loaded.state_dict().items():
assert torch.equal(value, model.state_dict()[name])
assert {id(parameter) for parameter in optimizer.param_groups[0]["params"]} == {
id(parameter) for parameter in loaded.parameters()
}


def test_training_script_resume_bypasses_fresh_model_factory(tmp_path, monkeypatch):
import yaml
from cosmodiff import optim, utils

config_path = tmp_path / "run.yaml"
config_path.write_text(yaml.safe_dump({"io": {"output_dir": str(tmp_path / "output")},
"train": {}, "data": {}}))
checkpoint = str(tmp_path / "checkpoint")
monkeypatch.setattr(utils, "find_latest_checkpoint", lambda _: checkpoint)
monkeypatch.setattr(utils, "parse_config_data", lambda _: {"data": object()})

def unexpected_fresh_model(_):
raise AssertionError("resume must not reconstruct/reinitialize the model")

monkeypatch.setattr(utils, "parse_config_model", unexpected_fresh_model)
calls = []

def fake_train(dataset, **kwargs):
calls.append(kwargs)
return {"metrics": {"epoch_loss": [0.1], "epoch_times": [1.0]}}

monkeypatch.setattr(optim, "train", fake_train)
script = Path(__file__).resolve().parents[2] / "scripts" / "cosmodiff_train.py"
spec = importlib.util.spec_from_file_location("cosmodiff_train_zero_init_test", script)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
monkeypatch.setattr(sys, "argv", [str(script), "--config", str(config_path)])
module.main()
assert len(calls) == 1
assert calls[0]["resume_from_checkpoint"] == checkpoint
15 changes: 15 additions & 0 deletions cosmodiff/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -458,6 +458,12 @@ def parse_config_model(config: dict):
parsed yaml config dict. Any missing keys return ``None``, which will
trigger the corresponding default in ``train()``.

``model.initialization`` defaults to ``"native"`` (unchanged diffusers
initialization). ``"adaln_zero"`` opts a fresh ``DiTTransformer2DModel``
into zero modulation/output initialization before the optimizer is built.
Checkpoint loading and the training script's resume path do not use this
fresh-model factory and never reapply initialization.

Args:
config (dict): Parsed yaml config, e.g. from ``yaml.safe_load()``.

Expand Down Expand Up @@ -486,8 +492,17 @@ def parse_config_model(config: dict):
# --- model ----------------------------------------------------------
model = None
if "model" in config:
initialization = config["model"].get("initialization", "native")
if initialization not in ("native", "adaln_zero"):
raise ValueError("model.initialization must be 'native' or 'adaln_zero'")
if (initialization == "adaln_zero"
and config["model"]["class"] != "DiTTransformer2DModel"):
raise ValueError("model.initialization='adaln_zero' requires DiTTransformer2DModel")
model_cls = getattr(diffusers, config["model"]["class"])
model = model_cls(**config["model"].get("kwargs", {}))
if initialization == "adaln_zero":
from .initialization import initialize_dit_adaln_zero
initialize_dit_adaln_zero(model, fresh=True)
model.to(device)

# --- optimizer ------------------------------------------------------
Expand Down
Loading