Skip to content
Merged
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
17 changes: 17 additions & 0 deletions miles/rollout/encoder_hub/__init__.py
Original file line number Diff line number Diff line change
@@ -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}")
58 changes: 58 additions & 0 deletions miles/rollout/encoder_hub/wan2_2.py
Original file line number Diff line number Diff line change
@@ -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,
}
Loading