# Copyright 2026 The HuggingFace Team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. from __future__ import annotations import argparse import math from collections.abc import Mapping from typing import Dict, Literal, Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_utils import ModelMixin DICO_PRESET_CONFIGS: Dict[str, Dict[str, object]] = { "DiCo-S": { "hidden_size": 128, "depth": [5, 4, 4, 4, 4], "mlp_ratio": 2.0, }, "DiCo-B": { "hidden_size": 256, "depth": [5, 4, 4, 4, 4], "mlp_ratio": 2.0, }, "DiCo-L": { "hidden_size": 352, "depth": [9, 8, 9, 8, 9], "mlp_ratio": 2.0, }, "DiCo-XL": { "hidden_size": 416, "depth": [9, 9, 10, 9, 9], "mlp_ratio": 2.0, }, "DiCo-H": { "hidden_size": 416, "depth": [14, 12, 10, 12, 14], "mlp_ratio": 4.0, }, } def remap_legacy_state_dict(state_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: """Map wrapper/backbone keys from legacy checkpoints to native model keys.""" remapped: Dict[str, torch.Tensor] = {} for key, value in state_dict.items(): new_key = key for prefix in ("transformer.", "model.", "net."): if new_key.startswith(prefix): new_key = new_key[len(prefix) :] break remapped[new_key] = value return remapped def infer_learn_sigma(state_dict: Dict[str, torch.Tensor], in_channels: int = 4) -> bool: weight = state_dict.get("final_layer.out_proj.weight") if weight is None: return True return int(weight.shape[0]) == in_channels * 2 def config_from_legacy(config: Dict[str, object]) -> Dict[str, object]: """Build native config kwargs from a legacy config.json dict.""" model_type = config.get("model_type") or config.get("model_name") or config.get("model") if model_type not in DICO_PRESET_CONFIGS: raise ValueError(f"Unknown DiCo preset '{model_type}'. Known: {list(DICO_PRESET_CONFIGS)}") preset = dict(DICO_PRESET_CONFIGS[model_type]) preset["num_classes"] = int(config.get("num_class_embeds") or config.get("num_classes") or 1000) preset["model_type"] = model_type preset["input_size"] = int(config.get("input_size") or config.get("sample_size") or 32) if config.get("learn_sigma") is not None: preset["learn_sigma"] = bool(config["learn_sigma"]) return preset def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: return x * (1 + scale.unsqueeze(-1).unsqueeze(-1)) + shift.unsqueeze(-1).unsqueeze(-1) class LayerNorm2d(nn.LayerNorm): def __init__(self, num_channels: int, eps: float = 1e-6, affine: bool = True): super().__init__(num_channels, eps=eps, elementwise_affine=affine) def forward(self, x: torch.Tensor) -> torch.Tensor: x = x.permute(0, 2, 3, 1) x = F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps) return x.permute(0, 3, 1, 2) class DiCoTimestepEmbedder(nn.Module): def __init__(self, hidden_size: int, frequency_embedding_size: int = 256): super().__init__() self.mlp = nn.Sequential( nn.Linear(frequency_embedding_size, hidden_size, bias=True), nn.SiLU(), nn.Linear(hidden_size, hidden_size, bias=True), ) self.frequency_embedding_size = frequency_embedding_size @staticmethod def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000) -> torch.Tensor: half = dim // 2 freqs = torch.exp( -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half ).to(device=t.device) args = t[:, None].float() * freqs[None] embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) if dim % 2: embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding def forward(self, t: torch.Tensor) -> torch.Tensor: t_freq = self.timestep_embedding(t, self.frequency_embedding_size) weight_dtype = self.mlp[0].weight.dtype return self.mlp(t_freq.to(dtype=weight_dtype)) class DiCoLabelEmbedder(nn.Module): def __init__(self, num_classes: int, hidden_size: int, dropout_prob: float): super().__init__() use_cfg_embedding = dropout_prob > 0 self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size) self.num_classes = num_classes self.dropout_prob = dropout_prob def token_drop(self, labels: torch.Tensor, force_drop_ids: Optional[torch.Tensor] = None) -> torch.Tensor: if force_drop_ids is None: drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob else: drop_ids = force_drop_ids == 1 return torch.where(drop_ids, self.num_classes, labels) def forward( self, labels: torch.Tensor, train: bool, force_drop_ids: Optional[torch.Tensor] = None, ) -> torch.Tensor: use_dropout = self.dropout_prob > 0 if (train and use_dropout) or (force_drop_ids is not None): labels = self.token_drop(labels, force_drop_ids) return self.embedding_table(labels) class DiCoMultiScaleLabelEmbedder(nn.Module): def __init__( self, num_classes: int, hidden_size_0: int, hidden_size_1: int, hidden_size_2: int, dropout_prob: float, ): super().__init__() use_cfg_embedding = dropout_prob > 0 self.embedding_table_0 = nn.Embedding(num_classes + use_cfg_embedding, hidden_size_0) self.embedding_table_1 = nn.Embedding(num_classes + use_cfg_embedding, hidden_size_1) self.embedding_table_2 = nn.Embedding(num_classes + use_cfg_embedding, hidden_size_2) self.num_classes = num_classes self.dropout_prob = dropout_prob def token_drop(self, labels: torch.Tensor, force_drop_ids: Optional[torch.Tensor] = None) -> torch.Tensor: if force_drop_ids is None: drop_ids = torch.rand(labels.shape[0], device=labels.device) < self.dropout_prob else: drop_ids = force_drop_ids == 1 return torch.where(drop_ids, self.num_classes, labels) def forward( self, labels: torch.Tensor, train: bool, force_drop_ids: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: use_dropout = self.dropout_prob > 0 if (train and use_dropout) or (force_drop_ids is not None): labels = self.token_drop(labels, force_drop_ids) return ( self.embedding_table_0(labels), self.embedding_table_1(labels), self.embedding_table_2(labels), ) class DiCoBlock(nn.Module): def __init__(self, hidden_size: int, mlp_ratio: float = 4.0): super().__init__() self.conv1 = nn.Conv2d(hidden_size, hidden_size, kernel_size=1, padding=0, stride=1, groups=1, bias=True) self.conv2 = nn.Conv2d( hidden_size, hidden_size, kernel_size=3, padding=1, stride=1, groups=hidden_size, bias=True ) self.conv3 = nn.Conv2d(hidden_size, hidden_size, kernel_size=1, padding=0, stride=1, groups=1, bias=True) self.ca = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(hidden_size, hidden_size, kernel_size=1, padding=0, stride=1, groups=1, bias=True), nn.Sigmoid(), ) ffn_channel = int(mlp_ratio * hidden_size) self.conv4 = nn.Conv2d(hidden_size, ffn_channel, kernel_size=1, padding=0, stride=1, groups=1, bias=True) self.conv5 = nn.Conv2d(ffn_channel, hidden_size, kernel_size=1, padding=0, stride=1, groups=1, bias=True) self.norm1 = LayerNorm2d(hidden_size, affine=False, eps=1e-6) self.norm2 = LayerNorm2d(hidden_size, affine=False, eps=1e-6) self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) def forward(self, inp: torch.Tensor, c: torch.Tensor) -> torch.Tensor: shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1) x = modulate(self.norm1(inp), shift_msa, scale_msa) x = F.gelu(self.conv2(self.conv1(x))) x = x * self.ca(x) x = self.conv3(x) x = inp + gate_msa.unsqueeze(-1).unsqueeze(-1) * x x = x + gate_mlp.unsqueeze(-1).unsqueeze(-1) * self.conv5( F.gelu(self.conv4(modulate(self.norm2(x), shift_mlp, scale_mlp))) ) return x class DiCoFinalLayer(nn.Module): def __init__(self, hidden_size: int, out_channels: int): super().__init__() self.norm_final = LayerNorm2d(hidden_size, affine=False, eps=1e-6) self.out_proj = nn.Conv2d(hidden_size, out_channels, kernel_size=3, stride=1, padding=1, bias=True) self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor: shift, scale = self.adaLN_modulation(c).chunk(2, dim=1) x = modulate(self.norm_final(x), shift, scale) return self.out_proj(x) class OverlapPatchEmbed(nn.Module): def __init__(self, in_c: int = 3, embed_dim: int = 48, bias: bool = False): super().__init__() self.proj = nn.Conv2d(in_c, embed_dim, kernel_size=3, stride=1, padding=1, bias=bias) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.proj(x) class Downsample(nn.Module): def __init__(self, n_feat: int): super().__init__() self.body = nn.Sequential( nn.Conv2d(n_feat, n_feat // 2, kernel_size=3, stride=1, padding=1, bias=False), nn.PixelUnshuffle(2), ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.body(x) class Upsample(nn.Module): def __init__(self, n_feat: int): super().__init__() self.body = nn.Sequential( nn.Conv2d(n_feat, n_feat * 2, kernel_size=3, stride=1, padding=1, bias=False), nn.PixelShuffle(2), ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.body(x) class DiCoTransformer2DModel(ModelMixin, ConfigMixin): r""" DiCo (Diffusion ConvNet) denoiser for class-conditional latent diffusion. ConvNet U-Net backbone with multi-scale adaLN conditioning, operating on VAE latents. """ _supports_gradient_checkpointing = True @register_to_config def __init__( self, input_size: int = 32, in_channels: int = 4, hidden_size: int = 416, depth: Optional[list[int]] = None, mlp_ratio: float = 2.0, class_dropout_prob: float = 0.1, num_classes: int = 1000, learn_sigma: bool = True, model_type: str | None = None, num_class_embeds: int | None = None, ): super().__init__() if num_class_embeds is not None: num_classes = int(num_class_embeds) if model_type in DICO_PRESET_CONFIGS: preset = DICO_PRESET_CONFIGS[model_type] hidden_size = int(preset["hidden_size"]) depth = list(preset["depth"]) mlp_ratio = float(preset["mlp_ratio"]) if depth is None: depth = [9, 9, 10, 9, 9] self.learn_sigma = learn_sigma self.in_channels = in_channels self.out_channels = in_channels * 2 if learn_sigma else in_channels self.num_classes = num_classes self.gradient_checkpointing = False self.x_embedder = OverlapPatchEmbed(in_channels, hidden_size, bias=True) self.t_embedder_1 = DiCoTimestepEmbedder(hidden_size) self.y_embedder = DiCoMultiScaleLabelEmbedder( num_classes, hidden_size, hidden_size * 2, hidden_size * 4, class_dropout_prob ) self.t_embedder_2 = DiCoTimestepEmbedder(hidden_size * 2) self.t_embedder_3 = DiCoTimestepEmbedder(hidden_size * 4) self.encoder_level_1 = nn.ModuleList([DiCoBlock(hidden_size, mlp_ratio) for _ in range(depth[0])]) self.down1_2 = Downsample(hidden_size) self.encoder_level_2 = nn.ModuleList([DiCoBlock(hidden_size * 2, mlp_ratio) for _ in range(depth[1])]) self.down2_3 = Downsample(hidden_size * 2) self.latent = nn.ModuleList([DiCoBlock(hidden_size * 4, mlp_ratio) for _ in range(depth[2])]) self.up3_2 = Upsample(int(hidden_size * 4)) self.reduce_chan_level2 = nn.Conv2d(int(hidden_size * 4), int(hidden_size * 2), kernel_size=1, bias=True) self.decoder_level_2 = nn.ModuleList([DiCoBlock(hidden_size * 2, mlp_ratio) for _ in range(depth[3])]) self.up2_1 = Upsample(int(hidden_size * 2)) self.reduce_chan_level1 = nn.Conv2d(int(hidden_size * 2), int(hidden_size * 2), kernel_size=1, bias=True) self.decoder_level_1 = nn.ModuleList([DiCoBlock(hidden_size * 2, mlp_ratio) for _ in range(depth[4])]) self.final_layer = DiCoFinalLayer(hidden_size * 2, self.out_channels) self.initialize_weights() def initialize_weights(self) -> None: def _basic_init(module: nn.Module): if isinstance(module, (nn.Linear, nn.Conv2d)): torch.nn.init.xavier_uniform_(module.weight) if module.bias is not None: nn.init.constant_(module.bias, 0) self.apply(_basic_init) w = self.x_embedder.proj.weight.data nn.init.xavier_uniform_(w.view([w.shape[0], -1])) nn.init.constant_(self.x_embedder.proj.bias, 0) nn.init.normal_(self.y_embedder.embedding_table_0.weight, std=0.02) nn.init.normal_(self.y_embedder.embedding_table_1.weight, std=0.02) nn.init.normal_(self.y_embedder.embedding_table_2.weight, std=0.02) for embedder in (self.t_embedder_1, self.t_embedder_2, self.t_embedder_3): nn.init.normal_(embedder.mlp[0].weight, std=0.02) nn.init.normal_(embedder.mlp[2].weight, std=0.02) blocks = ( self.encoder_level_1 + self.encoder_level_2 + self.latent + self.decoder_level_2 + self.decoder_level_1 ) for block in blocks: nn.init.constant_(block.adaLN_modulation[-1].weight, 0) nn.init.constant_(block.adaLN_modulation[-1].bias, 0) nn.init.constant_(self.final_layer.adaLN_modulation[-1].weight, 0) nn.init.constant_(self.final_layer.adaLN_modulation[-1].bias, 0) nn.init.constant_(self.final_layer.out_proj.weight, 0) nn.init.constant_(self.final_layer.out_proj.bias, 0) def _run_block(self, block: DiCoBlock, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor: if self.training and self.gradient_checkpointing: return torch.utils.checkpoint.checkpoint(block, x, c, use_reentrant=False) return block(x, c) def forward( self, hidden_states: torch.Tensor, timestep: torch.LongTensor, class_labels: torch.LongTensor, force_drop_ids: Optional[torch.Tensor] = None, return_dict: bool = True, ) -> Transformer2DModelOutput | Tuple: timestep = torch.as_tensor(timestep, device=hidden_states.device) if timestep.ndim == 0: timestep = timestep.repeat(hidden_states.shape[0]) else: timestep = timestep.reshape(-1) if timestep.shape[0] == 1 and hidden_states.shape[0] > 1: timestep = timestep.repeat(hidden_states.shape[0]) x = self.x_embedder(hidden_states) t1 = self.t_embedder_1(timestep) y1, y2, y3 = self.y_embedder(class_labels, self.training, force_drop_ids=force_drop_ids) c1 = t1 + y1 c2 = self.t_embedder_2(timestep) + y2 c3 = self.t_embedder_3(timestep) + y3 out_enc_level1 = x for block in self.encoder_level_1: out_enc_level1 = self._run_block(block, out_enc_level1, c1) out_enc_level2 = self.down1_2(out_enc_level1) for block in self.encoder_level_2: out_enc_level2 = self._run_block(block, out_enc_level2, c2) latent = self.down2_3(out_enc_level2) for block in self.latent: latent = self._run_block(block, latent, c3) inp_dec_level2 = self.reduce_chan_level2(torch.cat([self.up3_2(latent), out_enc_level2], dim=1)) for block in self.decoder_level_2: inp_dec_level2 = self._run_block(block, inp_dec_level2, c2) inp_dec_level1 = self.reduce_chan_level1(torch.cat([self.up2_1(inp_dec_level2), out_enc_level1], dim=1)) for block in self.decoder_level_1: inp_dec_level1 = self._run_block(block, inp_dec_level1, c2) output = self.final_layer(inp_dec_level1, c2) if not return_dict: return (output,) return Transformer2DModelOutput(sample=output) @classmethod def from_dico_checkpoint( cls, checkpoint_path: str, weights: Literal["model", "ema"] = "ema", map_location: str = "cpu", strict: bool = True, model_type: str | None = None, ) -> Tuple["DiCoTransformer2DModel", Dict[str, object]]: checkpoint = torch.load(checkpoint_path, map_location=map_location, weights_only=False) state_dict = checkpoint if isinstance(checkpoint, Mapping): if weights in checkpoint: state_dict = checkpoint[weights] elif "state_dict" in checkpoint: state_dict = checkpoint["state_dict"] state_dict = remap_legacy_state_dict(state_dict) ckpt_args = checkpoint.get("args") if isinstance(checkpoint, Mapping) else None args_dict: Dict[str, object] = {} if ckpt_args is not None: if isinstance(ckpt_args, argparse.Namespace): args_dict = vars(ckpt_args) elif isinstance(ckpt_args, Mapping): args_dict = dict(ckpt_args) resolved_model_type = model_type or args_dict.get("model") or args_dict.get("model_type") image_size = int(args_dict.get("image_size") or 256) num_classes = int(args_dict.get("num_classes") or 1000) config: Dict[str, object] = { "input_size": image_size // 8, "num_classes": num_classes, "learn_sigma": infer_learn_sigma(state_dict), } if resolved_model_type in DICO_PRESET_CONFIGS: config["model_type"] = resolved_model_type model = cls(**config) model.load_state_dict(state_dict, strict=strict) metadata = { "checkpoint_path": checkpoint_path, "weights": weights, "model_type": resolved_model_type, "source_args": ckpt_args, } return model, metadata DiCoDiffusersModel = DiCoTransformer2DModel