from __future__ import annotations import math from dataclasses import dataclass from math import sqrt from typing import Dict, 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_utils import ModelMixin from diffusers.models.normalization import RMSNorm from diffusers.utils import BaseOutput PMF_PRESET_CONFIGS: Dict[str, Dict[str, object]] = { "pMF-B/16": { "sample_size": 256, "patch_size": 16, "hidden_size": 768, "depth": 16, "num_attention_heads": 12, "bottleneck_dim": 128, "aux_head_depth": 8, }, "pMF-B/32": { "sample_size": 512, "patch_size": 32, "hidden_size": 768, "depth": 16, "num_attention_heads": 12, "bottleneck_dim": 128, "aux_head_depth": 8, }, "pMF-L/16": { "sample_size": 256, "patch_size": 16, "hidden_size": 1024, "depth": 32, "num_attention_heads": 16, "bottleneck_dim": 128, "aux_head_depth": 8, }, "pMF-L/32": { "sample_size": 512, "patch_size": 32, "hidden_size": 1024, "depth": 32, "num_attention_heads": 16, "bottleneck_dim": 128, "aux_head_depth": 8, }, "pMF-H/16": { "sample_size": 256, "patch_size": 16, "hidden_size": 1280, "depth": 48, "num_attention_heads": 16, "bottleneck_dim": 256, "aux_head_depth": 8, }, "pMF-H/32": { "sample_size": 512, "patch_size": 32, "hidden_size": 1280, "depth": 48, "num_attention_heads": 16, "bottleneck_dim": 256, "aux_head_depth": 8, }, } RECOMMENDED_NOISE_BY_MODEL: Dict[str, float] = { "pMF-B/16": 1.0, "pMF-B/32": 2.0, "pMF-L/16": 1.0, "pMF-L/32": 4.0, "pMF-H/16": 2.0, "pMF-H/32": 4.0, } # Legacy torch repo keys (pmfDiT_*) LEGACY_MODEL_ALIASES: Dict[str, str] = { "pmfDiT_B_16": "pMF-B/16", "pmfDiT_B_32": "pMF-B/32", "pmfDiT_L_16": "pMF-L/16", "pmfDiT_L_32": "pMF-L/32", "pmfDiT_H_16": "pMF-H/16", "pmfDiT_H_32": "pMF-H/32", } @dataclass class PMFTransformer2DOutput(BaseOutput): u: torch.Tensor v: Optional[torch.Tensor] = None def remap_legacy_state_dict(state_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: """Map wrapper/backbone keys from legacy checkpoints to native PMFTransformer2DModel keys.""" remapped: Dict[str, torch.Tensor] = {} for key, value in state_dict.items(): new_key = key for prefix in ("transformer.", "net."): if new_key.startswith(prefix): new_key = new_key[len(prefix) :] break # Official PyTorch checkpoints use TorchLinear/TorchEmbedding wrappers. new_key = new_key.replace("._flax_linear", "").replace("._flax_embedding", "") if new_key == "rope_freqs": continue remapped[new_key] = value return remapped 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_str") if model_type in LEGACY_MODEL_ALIASES: model_type = LEGACY_MODEL_ALIASES[model_type] if model_type not in PMF_PRESET_CONFIGS: raise ValueError(f"Unknown pMF preset '{model_type}'. Known: {list(PMF_PRESET_CONFIGS)}") preset = dict(PMF_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 if config.get("sample_size") is not None: preset["sample_size"] = int(config["sample_size"]) if config.get("eval_mode") is not None: preset["eval_mode"] = bool(config["eval_mode"]) return preset def _scaled_linear( in_features: int, out_features: int, *, bias: bool = True, weight_init: str = "scaled_variance", init_constant: float = 1.0, bias_init: str = "zeros", ) -> nn.Linear: layer = nn.Linear(in_features, out_features, bias=bias) if weight_init == "scaled_variance": std = init_constant / sqrt(in_features) nn.init.normal_(layer.weight, std=std) elif weight_init == "zeros": nn.init.zeros_(layer.weight) else: raise ValueError(f"Invalid weight_init: {weight_init}") if bias: if bias_init == "zeros": nn.init.zeros_(layer.bias) else: raise ValueError(f"Invalid bias_init: {bias_init}") return layer class PMFTimestepEmbedder(nn.Module): def __init__( self, hidden_size: int, frequency_embedding_size: int = 256, init_constant: float = 1.0, ): super().__init__() init_kwargs = dict( out_features=hidden_size, bias=True, weight_init="scaled_variance", init_constant=init_constant, bias_init="zeros", ) self.mlp = nn.Sequential( _scaled_linear(frequency_embedding_size, **init_kwargs), nn.SiLU(), _scaled_linear(hidden_size, **init_kwargs), ) 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, device=t.device) / half ) 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) return self.mlp(t_freq) class PMFLabelEmbedder(nn.Module): def __init__(self, num_classes: int, hidden_size: int, init_constant: float = 1.0): super().__init__() self.embedding_table = nn.Embedding(num_classes + 1, hidden_size) nn.init.normal_(self.embedding_table.weight, std=init_constant / sqrt(hidden_size)) def forward(self, labels: torch.Tensor) -> torch.Tensor: return self.embedding_table(labels) class PMFBottleneckPatchEmbedder(nn.Module): def __init__( self, input_size: int, patch_size: int, pca_channels: int, in_channels: int, hidden_size: int, bias: bool = True, ): super().__init__() self.patch_size = (patch_size, patch_size) self.num_patches = (input_size // patch_size) ** 2 self.proj1 = nn.Conv2d( in_channels, pca_channels, kernel_size=patch_size, stride=patch_size, bias=bias, ) self.proj2 = nn.Conv2d(pca_channels, hidden_size, kernel_size=1, stride=1, bias=bias) kh = kw = patch_size fan_in = kh * kw * in_channels fan_out = pca_channels limit = math.sqrt(6.0 / (fan_in + fan_out)) nn.init.uniform_(self.proj1.weight, -limit, limit) fan_in = pca_channels fan_out = hidden_size limit = math.sqrt(6.0 / (fan_in + fan_out)) nn.init.uniform_(self.proj2.weight, -limit, limit) if bias: nn.init.zeros_(self.proj1.bias) nn.init.zeros_(self.proj2.bias) def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.proj2(self.proj1(x)) return x.flatten(2).transpose(1, 2) def precompute_rope_freqs(dim: int, seq_len: int, theta: float = 10000.0) -> torch.Tensor: dim = dim // 2 grid_size = int(seq_len**0.5) freqs = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) positions = torch.arange(grid_size, dtype=torch.float32) freqs_h = torch.einsum("i,j->ij", positions, freqs) freqs_w = torch.einsum("i,j->ij", positions, freqs) freqs_2d = torch.cat( [ torch.tile(freqs_h[:, None, :], (1, grid_size, 1)), torch.tile(freqs_w[None, :, :], (grid_size, 1, 1)), ], dim=-1, ) real = torch.cos(freqs_2d).reshape(seq_len, dim) imag = torch.sin(freqs_2d).reshape(seq_len, dim) return torch.complex(real, imag) def apply_rotary_pos_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor: x_float = x.to(torch.float32) x_complex = torch.view_as_complex(x_float.reshape(*x_float.shape[:-1], -1, 2).contiguous()) freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(2) token_count = freqs_cis.shape[1] x_rotated = x_complex.clone() x_rotated[:, -token_count:, :] = x_complex[:, -token_count:, :] * freqs_cis x_out = torch.view_as_real(x_rotated).flatten(-2) return x_out.to(x.dtype) class PMFAttention(nn.Module): def __init__( self, hidden_size: int, num_heads: int, weight_init_constant: float = 0.32, eps: float = 1e-6, ): super().__init__() self.num_heads = num_heads self.head_dim = hidden_size // num_heads init_kwargs = dict( bias=False, weight_init="scaled_variance", init_constant=weight_init_constant, ) self.q_proj = _scaled_linear(hidden_size, hidden_size, **init_kwargs) self.k_proj = _scaled_linear(hidden_size, hidden_size, **init_kwargs) self.v_proj = _scaled_linear(hidden_size, hidden_size, **init_kwargs) self.out_proj = _scaled_linear(hidden_size, hidden_size, **init_kwargs) self.q_norm = RMSNorm(self.head_dim, eps=eps) self.k_norm = RMSNorm(self.head_dim, eps=eps) def forward(self, x: torch.Tensor, rope_freqs: torch.Tensor) -> torch.Tensor: batch_size, seq_len, channels = x.shape q = self.q_proj(x).reshape(batch_size, seq_len, self.num_heads, self.head_dim) k = self.k_proj(x).reshape(batch_size, seq_len, self.num_heads, self.head_dim) v = self.v_proj(x).reshape(batch_size, seq_len, self.num_heads, self.head_dim) q = self.q_norm(q) k = self.k_norm(k) q = apply_rotary_pos_emb(q, rope_freqs) k = apply_rotary_pos_emb(k, rope_freqs) query = q / math.sqrt(self.head_dim) attn_weights = torch.einsum("bqhd,bkhd->bhqk", query, k) attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32) attn = torch.einsum("bhqk,bkhd->bqhd", attn_weights, v) attn = attn.reshape(batch_size, seq_len, channels) return self.out_proj(attn) class PMFSwiGLUMlp(nn.Module): def __init__(self, dim: int, hidden_dim: int, weight_init_constant: float = 0.32): super().__init__() init_kwargs = dict(bias=False, weight_init="scaled_variance", init_constant=weight_init_constant) self.w1 = _scaled_linear(dim, hidden_dim, **init_kwargs) self.w3 = _scaled_linear(dim, hidden_dim, **init_kwargs) self.w2 = _scaled_linear(hidden_dim, dim, **init_kwargs) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.w2(F.silu(self.w1(x)) * self.w3(x)) class PMFTransformerBlock(nn.Module): def __init__( self, hidden_size: int, num_heads: int, mlp_ratio: float = 8 / 3, weight_init_constant: float = 0.32, eps: float = 1e-6, ): super().__init__() self.norm1 = RMSNorm(hidden_size, eps=eps) self.attn = PMFAttention(hidden_size, num_heads, weight_init_constant=weight_init_constant, eps=eps) self.norm2 = RMSNorm(hidden_size, eps=eps) mlp_hidden_dim = int(hidden_size * mlp_ratio) if hidden_size > 1024: mlp_hidden_dim = (mlp_hidden_dim + 7) // 8 * 8 self.mlp = PMFSwiGLUMlp(hidden_size, mlp_hidden_dim, weight_init_constant=weight_init_constant) self.attn_scale = nn.Parameter(torch.zeros(hidden_size)) self.mlp_scale = nn.Parameter(torch.zeros(hidden_size)) def forward(self, x: torch.Tensor, rope_freqs: torch.Tensor) -> torch.Tensor: x = x + self.attn(self.norm1(x), rope_freqs) * self.attn_scale x = x + self.mlp(self.norm2(x)) * self.mlp_scale return x class PMFFinalLayer(nn.Module): def __init__(self, hidden_size: int, patch_size: int, out_channels: int, eps: float = 1e-6): super().__init__() self.norm = RMSNorm(hidden_size, eps=eps) self.linear = _scaled_linear( hidden_size, patch_size * patch_size * out_channels, bias=True, weight_init="zeros", bias_init="zeros", ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.linear(self.norm(x)) class PMFTransformer2DModel(ModelMixin, ConfigMixin): """Native diffusers implementation of the pMF DiT backbone.""" _supports_gradient_checkpointing = True _skip_layerwise_casting_patterns = ["pos_embed", "rope_freqs"] @register_to_config def __init__( self, sample_size: int = 256, patch_size: int = 16, in_channels: int = 3, hidden_size: int = 768, depth: int = 16, num_attention_heads: int = 12, mlp_ratio: float = 8 / 3, num_classes: int = 1000, bottleneck_dim: int = 128, aux_head_depth: int = 8, num_class_tokens: int = 8, num_time_tokens: int = 4, num_cfg_tokens: int = 4, num_interval_tokens: int = 2, token_init_constant: float = 1.0, embedding_init_constant: float = 1.0, weight_init_constant: float = 0.32, eval_mode: bool = True, model_type: str | None = None, num_class_embeds: int | None = None, t_clip_min: float = 0.05, norm_eps: float = 1e-6, ): super().__init__() if num_class_embeds is not None: num_classes = int(num_class_embeds) if model_type in LEGACY_MODEL_ALIASES: model_type = LEGACY_MODEL_ALIASES[model_type] if model_type in PMF_PRESET_CONFIGS: preset = PMF_PRESET_CONFIGS[model_type] sample_size = int(preset["sample_size"]) patch_size = int(preset["patch_size"]) hidden_size = int(preset["hidden_size"]) depth = int(preset["depth"]) num_attention_heads = int(preset["num_attention_heads"]) bottleneck_dim = int(preset["bottleneck_dim"]) aux_head_depth = int(preset["aux_head_depth"]) self.sample_size = sample_size self.patch_size = patch_size self.in_channels = in_channels self.out_channels = in_channels self.hidden_size = hidden_size self.depth = depth self.num_attention_heads = num_attention_heads self.aux_head_depth = aux_head_depth self.num_class_tokens = num_class_tokens self.num_time_tokens = num_time_tokens self.num_cfg_tokens = num_cfg_tokens self.num_interval_tokens = num_interval_tokens self.prefix_tokens = ( num_class_tokens + num_cfg_tokens + 2 * num_interval_tokens + num_time_tokens ) self.t_clip_min = t_clip_min self.eval_mode = eval_mode self.gradient_checkpointing = False self.x_embedder = PMFBottleneckPatchEmbedder( sample_size, patch_size, bottleneck_dim, in_channels, hidden_size, bias=True, ) embed_kwargs = dict(hidden_size=hidden_size, init_constant=embedding_init_constant) self.h_embedder = PMFTimestepEmbedder(**embed_kwargs) self.omega_embedder = PMFTimestepEmbedder(**embed_kwargs) self.cfg_t_start_embedder = PMFTimestepEmbedder(**embed_kwargs) self.cfg_t_end_embedder = PMFTimestepEmbedder(**embed_kwargs) self.y_embedder = PMFLabelEmbedder(num_classes, hidden_size, init_constant=embedding_init_constant) token_std = token_init_constant / math.sqrt(hidden_size) self.time_tokens = nn.Parameter(torch.randn(1, num_time_tokens, hidden_size) * token_std) self.class_tokens = nn.Parameter(torch.randn(1, num_class_tokens, hidden_size) * token_std) self.omega_tokens = nn.Parameter(torch.randn(1, num_cfg_tokens, hidden_size) * token_std) self.t_min_tokens = nn.Parameter(torch.randn(1, num_interval_tokens, hidden_size) * token_std) self.t_max_tokens = nn.Parameter(torch.randn(1, num_interval_tokens, hidden_size) * token_std) total_tokens = self.x_embedder.num_patches + self.prefix_tokens self.pos_embed = nn.Parameter(torch.randn(1, total_tokens, hidden_size) * 0.02) head_dim = hidden_size // num_attention_heads self.register_buffer( "rope_freqs", precompute_rope_freqs(head_dim, self.x_embedder.num_patches), persistent=False, ) shared_depth = depth - aux_head_depth block_kwargs = dict( hidden_size=hidden_size, num_heads=num_attention_heads, mlp_ratio=mlp_ratio, weight_init_constant=weight_init_constant, eps=norm_eps, ) self.shared_blocks = nn.ModuleList([PMFTransformerBlock(**block_kwargs) for _ in range(shared_depth)]) self.u_heads = nn.ModuleList([PMFTransformerBlock(**block_kwargs) for _ in range(aux_head_depth)]) self.v_heads = nn.ModuleList( [PMFTransformerBlock(**block_kwargs) for _ in range(aux_head_depth if not eval_mode else 0)] ) self.u_final_layer = PMFFinalLayer(hidden_size, patch_size, in_channels, eps=norm_eps) self.v_final_layer = ( PMFFinalLayer(hidden_size, patch_size, in_channels, eps=norm_eps) if not eval_mode else None ) def _build_sequence( self, sample: torch.Tensor, h: torch.Tensor, omega: torch.Tensor, t_min: torch.Tensor, t_max: torch.Tensor, class_labels: torch.Tensor, ) -> torch.Tensor: x_embed = self.x_embedder(sample) h_embed = self.h_embedder(h) omega_embed = self.omega_embedder(1 - 1 / omega) t_min_embed = self.cfg_t_start_embedder(t_min) t_max_embed = self.cfg_t_end_embedder(t_max) y_embed = self.y_embedder(class_labels) time_tokens = self.time_tokens + h_embed.unsqueeze(1) omega_tokens = self.omega_tokens + omega_embed.unsqueeze(1) t_min_tokens = self.t_min_tokens + t_min_embed.unsqueeze(1) t_max_tokens = self.t_max_tokens + t_max_embed.unsqueeze(1) class_tokens = self.class_tokens + y_embed.unsqueeze(1) seq = torch.cat( [class_tokens, omega_tokens, t_min_tokens, t_max_tokens, time_tokens, x_embed], dim=1, ) return seq + self.pos_embed def _unpatchify(self, tokens: torch.Tensor) -> torch.Tensor: batch_size = tokens.shape[0] patch = self.patch_size grid = int(tokens.shape[1] ** 0.5) channels = self.out_channels x = tokens.reshape(batch_size, grid, grid, patch, patch, channels) x = torch.einsum("nhwpqc->nchpwq", x) return x.reshape(batch_size, channels, grid * patch, grid * patch) def forward( self, sample: torch.Tensor, timestep: torch.Tensor, class_labels: torch.Tensor, h: Optional[torch.Tensor] = None, omega: Optional[torch.Tensor] = None, guidance_interval_min: Optional[torch.Tensor] = None, guidance_interval_max: Optional[torch.Tensor] = None, return_dict: bool = True, ) -> PMFTransformer2DOutput | Tuple[torch.Tensor, Optional[torch.Tensor]]: batch_size = sample.shape[0] timestep = self._expand_batch(timestep, batch_size, sample.device, sample.dtype) h = self._expand_batch(h if h is not None else timestep, batch_size, sample.device, sample.dtype) omega = self._expand_batch( omega if omega is not None else torch.ones(batch_size, device=sample.device), batch_size, sample.device, sample.dtype, ) guidance_interval_min = self._expand_batch( guidance_interval_min if guidance_interval_min is not None else torch.zeros(batch_size, device=sample.device), batch_size, sample.device, sample.dtype, ) guidance_interval_max = self._expand_batch( guidance_interval_max if guidance_interval_max is not None else torch.ones(batch_size, device=sample.device), batch_size, sample.device, sample.dtype, ) seq = self._build_sequence(sample, h, omega, guidance_interval_min, guidance_interval_max, class_labels) rope_freqs = self.rope_freqs.to(device=sample.device) for block in self.shared_blocks: if self.training and self.gradient_checkpointing: seq = torch.utils.checkpoint.checkpoint(block, seq, rope_freqs, use_reentrant=False) else: seq = block(seq, rope_freqs) u_seq = v_seq = seq for block in self.u_heads: if self.training and self.gradient_checkpointing: u_seq = torch.utils.checkpoint.checkpoint(block, u_seq, rope_freqs, use_reentrant=False) else: u_seq = block(u_seq, rope_freqs) for block in self.v_heads: if self.training and self.gradient_checkpointing: v_seq = torch.utils.checkpoint.checkpoint(block, v_seq, rope_freqs, use_reentrant=False) else: v_seq = block(v_seq, rope_freqs) u_tokens = u_seq[:, self.prefix_tokens :] u_pred = self._unpatchify(self.u_final_layer(u_tokens)) t = timestep.reshape(batch_size, 1, 1, 1) u = (sample - u_pred) / torch.clamp(t, min=self.t_clip_min) v = None if self.v_final_layer is not None: v_tokens = v_seq[:, self.prefix_tokens :] v_pred = self._unpatchify(self.v_final_layer(v_tokens)) v = (sample - v_pred) / torch.clamp(t, min=self.t_clip_min) if not return_dict: return (u, v) return PMFTransformer2DOutput(u=u, v=v) @staticmethod def _expand_batch( value: torch.Tensor, batch_size: int, device: torch.device, dtype: torch.dtype, ) -> torch.Tensor: value = torch.as_tensor(value, device=device, dtype=dtype) if value.ndim == 0: value = value.reshape(1) if value.shape[0] == 1 and batch_size > 1: value = value.expand(batch_size) return value.reshape(batch_size) @classmethod def from_pmf_checkpoint( cls, checkpoint_path: str, model_type: str | None = None, map_location: str = "cpu", strict: bool = False, ) -> Tuple["PMFTransformer2DModel", Dict[str, object]]: checkpoint = torch.load(checkpoint_path, map_location=map_location, weights_only=False) if isinstance(checkpoint, dict) and "state_dict" in checkpoint: state_dict = checkpoint["state_dict"] else: state_dict = checkpoint if model_type is None: for key in ("model_type", "model_str", "model"): if isinstance(checkpoint, dict) and key in checkpoint: model_type = checkpoint[key] break if model_type in LEGACY_MODEL_ALIASES: model_type = LEGACY_MODEL_ALIASES[model_type] if model_type is None: raise ValueError("model_type is required when it cannot be inferred from the checkpoint.") config = dict(PMF_PRESET_CONFIGS[model_type]) config["model_type"] = model_type config["eval_mode"] = True model = cls(**config) model.load_state_dict(remap_legacy_state_dict(state_dict), strict=strict) metadata = {"checkpoint_path": checkpoint_path, "model_type": model_type} return model, metadata def to_pmf_checkpoint(self, prefix: str = "net.") -> Dict[str, torch.Tensor]: state_dict: Dict[str, torch.Tensor] = {} for key, value in self.state_dict().items(): state_dict[f"{prefix}{key}"] = value.detach().cpu() return state_dict @property def net(self): return self PMFDiffusersModel = PMFTransformer2DModel