diff --git a/miles/rollout/encoder_hub/__init__.py b/miles/rollout/encoder_hub/__init__.py new file mode 100644 index 00000000..4e04b245 --- /dev/null +++ b/miles/rollout/encoder_hub/__init__.py @@ -0,0 +1,17 @@ +"""Frozen-encoder logic per model family, decoupled from the training-side TrainPipelineConfig. + +Each family module provides: +- ``load_encoder(args, device)``: load the frozen encode components (tokenizer/text + encoder/VAE) from the ``--sft-encoder-checkpoint`` HF name or path; +- ``encode_sample(encoder, pixels, prompt, generator)``: encode one media/prompt pair + into a cached train sample (clean latent + cond kwargs); +- ``validate_args(args)``: family-specific encode constraints. +""" + + +def get_encoder(family: str | None): + if family == "wan2_2": + from miles.rollout.encoder_hub import wan2_2 + + return wan2_2 + raise ValueError(f"no encoder_hub entry for model family {family!r}") diff --git a/miles/rollout/encoder_hub/wan2_2.py b/miles/rollout/encoder_hub/wan2_2.py new file mode 100644 index 00000000..6a587c89 --- /dev/null +++ b/miles/rollout/encoder_hub/wan2_2.py @@ -0,0 +1,58 @@ +"""Wan2.2 frozen encoders (UMT5 text encoder + VAE) for offline SFT encoding.""" + +from __future__ import annotations + +import torch + + +def validate_args(args) -> None: + if (args.diffusion_output_num_frames - 1) % 4 != 0: + raise ValueError("--diffusion-output-num-frames must be 4k+1 for the Wan VAE temporal stride") + + +def load_encoder(args, device: torch.device) -> dict: + from diffusers import AutoencoderKLWan + from transformers import AutoTokenizer, UMT5EncoderModel + + ckpt = args.sft_encoder_checkpoint + tokenizer = AutoTokenizer.from_pretrained(ckpt, subfolder="tokenizer") + text_encoder = UMT5EncoderModel.from_pretrained(ckpt, subfolder="text_encoder", torch_dtype=torch.float32).to( + device + ) + vae = AutoencoderKLWan.from_pretrained(ckpt, subfolder="vae", torch_dtype=torch.float32).to(device) + view = (1, vae.config.z_dim, 1, 1, 1) + return { + "device": device, + "tokenizer": tokenizer, + "text_encoder": text_encoder, + "vae": vae, + "latents_mean": torch.tensor(vae.config.latents_mean).view(view).to(device), + "latents_std": torch.tensor(vae.config.latents_std).view(view).to(device), + } + + +@torch.no_grad() +def encode_sample(encoder: dict, pixels: torch.Tensor, prompt: str, generator: torch.Generator) -> dict: + from diffusers.pipelines.wan.pipeline_wan import prompt_clean + + device = encoder["device"] + latent = encoder["vae"].encode(pixels.unsqueeze(0).to(device, torch.float32)).latent_dist.sample(generator) + latent = (latent - encoder["latents_mean"]) / encoder["latents_std"] + + inputs = encoder["tokenizer"]( + [prompt_clean(prompt)], + padding="max_length", + max_length=512, + truncation=True, + add_special_tokens=True, + return_attention_mask=True, + return_tensors="pt", + ) + embeds = encoder["text_encoder"](inputs.input_ids.to(device), inputs.attention_mask.to(device)).last_hidden_state + embeds[:, int(inputs.attention_mask[0].sum()) :] = 0 + + return { + "latent": latent[0].to(torch.bfloat16).cpu(), + "cond_kwargs": {"encoder_hidden_states": embeds.to(torch.bfloat16).cpu()}, + "prompt": prompt, + }