Text-to-Image
MLX
Safetensors
Diffusion Single File
Anima-mlx / anima_mlx /config.py
fukujusou's picture
Upload folder using huggingface_hub
3bdc93d verified
Raw
History Blame Contribute Delete
8.43 kB
"""Frozen source-side model contract for the Anima MLX conversion.
This module intentionally records only facts confirmed from local artifacts.
Fields that require the original PyTorch pipeline, tokenizer, or scheduler are
left as ``None`` instead of guessed.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional
@dataclass(frozen=True)
class SafetensorsArtifact:
name: str
path: str
tensor_count: int
dtype: str
payload_bytes: int
file_size_bytes: int
@dataclass(frozen=True)
class TextEncoderContract:
artifact: SafetensorsArtifact
vocab_size: int
hidden_size: int
layer_count: int
mlp_intermediate_size: int
q_proj_out_features: int
kv_proj_out_features: int
o_proj_in_features: int
qk_norm_size: int
tokenizer_type: Optional[str] = None
tokenizer_vocab_path: Optional[str] = None
chat_template: Optional[str] = None
max_sequence_length: Optional[int] = None
source_max_length: Optional[int] = None
padding_side: Optional[str] = None
truncation_side: Optional[str] = None
pad_token_id: Optional[int] = None
bos_token_id: Optional[int] = None
eos_token_id: Optional[int] = None
rope_theta: Optional[float] = None
output_hidden_state: Optional[str] = None
pooling_or_norm: Optional[str] = None
uses_auxiliary_t5_token_ids: bool = False
@dataclass(frozen=True)
class DiffusionContract:
artifact: SafetensorsArtifact
dit_block_count: int
hidden_size: int
attention_head_dim: int
estimated_attention_heads: int
cross_attention_context_dim: int
x_embedder_input_dim: int
final_patch_dim: int
timestep_embedding_dim: int
adaln_hidden_dim: int
adaln_output_dim_per_sublayer: int
llm_adapter_block_count: int
llm_adapter_vocab_size: int
llm_adapter_hidden_size: int
llm_adapter_mlp_intermediate_size: int
latent_token_shape: Optional[str] = None
x_embedder_input_semantics: Optional[str] = None
input_latent_channels: Optional[int] = None
padding_mask_channels: Optional[int] = None
output_latent_channels: Optional[int] = None
patch_spatial: Optional[int] = None
patch_temporal: Optional[int] = None
max_img_h: Optional[int] = None
max_img_w: Optional[int] = None
max_frames: Optional[int] = None
pos_emb_cls: Optional[str] = None
pos_emb_learnable: Optional[bool] = None
pos_emb_interpolation: Optional[str] = None
rope_h_extrapolation_ratio: Optional[float] = None
rope_w_extrapolation_ratio: Optional[float] = None
rope_t_extrapolation_ratio: Optional[float] = None
min_fps: Optional[int] = None
max_fps: Optional[int] = None
timestep_range: Optional[str] = None
timestep_embedding_method: Optional[str] = None
adaln_split_order: Optional[str] = None
final_layer_adaln_split_order: Optional[str] = None
prediction_target: Optional[str] = None
scheduler_type: Optional[str] = None
scheduler_shift: Optional[float] = None
scheduler_multiplier: Optional[float] = None
cfg_formula: Optional[str] = None
@dataclass(frozen=True)
class VAEContract:
artifact: SafetensorsArtifact
encoder_input_channels: int
encoder_head_channels: int
decoder_input_channels: int
decoder_output_channels: int
base_channels: int
max_channels: int
middle_attention_channels: int
latent_scaling_factor: Optional[float] = None
latent_layout: Optional[str] = None
output_layout: Optional[str] = None
output_range: Optional[str] = None
temporal_downsample_ratio: Optional[int] = None
spatial_downsample_ratio: Optional[int] = None
norm_type: Optional[str] = None
attention_axes: Optional[str] = None
source_vae_class: Optional[str] = None
@dataclass(frozen=True)
class SourceContract:
text_encoder: TextEncoderContract
diffusion: DiffusionContract
vae: VAEContract
source_pipeline_path: Optional[str] = None
source_commit_or_version: Optional[str] = None
scheduler_config_path: Optional[str] = None
tokenizer_config_path: Optional[str] = None
TEXT_ENCODER_ARTIFACT = SafetensorsArtifact(
name="text_encoder",
path="split_files/text_encoders/qwen_3_06b_base-mlx.safetensors",
tensor_count=310,
dtype="BF16",
payload_bytes=1_192_099_840,
file_size_bytes=1_192_135_180,
)
DIFFUSION_ARTIFACT = SafetensorsArtifact(
name="diffusion",
path="split_files/diffusion_models/anima-base-v1.0-mlx.safetensors",
tensor_count=685,
dtype="BF16",
payload_bytes=4_182_137_856,
file_size_bytes=4_182_218_400,
)
VAE_ARTIFACT = SafetensorsArtifact(
name="vae",
path="split_files/vae/qwen_image_vae-mlx.safetensors",
tensor_count=108,
dtype="BF16",
payload_bytes=146_590_360,
file_size_bytes=146_603_060,
)
DEFAULT_SOURCE_CONTRACT = SourceContract(
text_encoder=TextEncoderContract(
artifact=TEXT_ENCODER_ARTIFACT,
vocab_size=151_936,
hidden_size=1_024,
layer_count=28,
mlp_intermediate_size=3_072,
q_proj_out_features=2_048,
kv_proj_out_features=1_024,
o_proj_in_features=2_048,
qk_norm_size=128,
tokenizer_type="ComfyUI AnimaTokenizer with Qwen2Tokenizer for qwen3_06b and T5TokenizerFast token IDs for t5xxl",
tokenizer_vocab_path="tokenizers/qwen25_tokenizer",
max_sequence_length=131_072,
source_max_length=99_999_999,
pad_token_id=151_643,
bos_token_id=151_643,
eos_token_id=151_645,
rope_theta=1_000_000.0,
output_hidden_state="last",
pooling_or_norm="ComfyUI Qwen3_06BModel uses layer_norm_hidden_state=False",
uses_auxiliary_t5_token_ids=True,
),
diffusion=DiffusionContract(
artifact=DIFFUSION_ARTIFACT,
dit_block_count=28,
hidden_size=2_048,
attention_head_dim=128,
estimated_attention_heads=16,
cross_attention_context_dim=1_024,
x_embedder_input_dim=68,
final_patch_dim=64,
timestep_embedding_dim=2_048,
adaln_hidden_dim=256,
adaln_output_dim_per_sublayer=6_144,
llm_adapter_block_count=6,
llm_adapter_vocab_size=32_128,
llm_adapter_hidden_size=1_024,
llm_adapter_mlp_intermediate_size=4_096,
latent_token_shape="B,C,T,H,W input -> B,T,H/2,W/2,D embedded patches",
x_embedder_input_semantics="(16 latent channels + 1 padding mask channel) * patch_temporal 1 * patch_spatial 2 * patch_spatial 2 = 68",
input_latent_channels=16,
padding_mask_channels=1,
output_latent_channels=16,
patch_spatial=2,
patch_temporal=1,
max_img_h=240,
max_img_w=240,
max_frames=128,
pos_emb_cls="rope3d",
pos_emb_learnable=True,
pos_emb_interpolation="crop",
rope_h_extrapolation_ratio=4.0,
rope_w_extrapolation_ratio=4.0,
rope_t_extrapolation_ratio=1.0,
min_fps=1,
max_fps=30,
timestep_range="ComfyUI ModelSamplingDiscreteFlow sigma in [0, 1], timestep = sigma * multiplier",
timestep_embedding_method="Cosmos Predict2 Timesteps sinusoidal embedding then TimestepEmbedding with AdaLN-LoRA",
adaln_split_order="shift, scale, gate",
final_layer_adaln_split_order="shift, scale",
prediction_target="flow/CONST denoising head as used by ComfyUI ModelType.FLOW",
scheduler_type="ComfyUI ModelSamplingDiscreteFlow",
scheduler_shift=3.0,
scheduler_multiplier=1.0,
),
vae=VAEContract(
artifact=VAE_ARTIFACT,
encoder_input_channels=3,
encoder_head_channels=32,
decoder_input_channels=16,
decoder_output_channels=3,
base_channels=96,
max_channels=384,
middle_attention_channels=384,
latent_layout="B,C,T,H,W",
output_layout="B,C,T,H,W before ComfyUI postprocess",
temporal_downsample_ratio=4,
spatial_downsample_ratio=8,
norm_type="WanVAE RMS_norm",
source_vae_class="comfy.ldm.wan.vae.WanVAE",
),
source_pipeline_path="bundled minimal MLX runtime",
source_commit_or_version="ComfyUI 25757a53c93281e8e2462ced8795373f09e675bf",
scheduler_config_path="anima_mlx/runtime/scheduler.py",
tokenizer_config_path="anima_mlx/runtime/tokenizer.py",
)