"""Direct two-block H3 text token refiner.""" import torch import torch.nn.functional as functional from torch import nn from .attention import rms_norm from .checkpoint import H3Checkpoint class _Linear(nn.Module): def __init__(self, checkpoint: H3Checkpoint, prefix: str, dtype: torch.dtype): super().__init__() self.register_buffer("weight", checkpoint.tensor(f"{prefix}.weight", dtype=dtype), persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return functional.linear(x, self.weight) class _RefinerBlock(nn.Module): def __init__(self, checkpoint: H3Checkpoint, prefix: str, dtype: torch.dtype): super().__init__() self.qkv = _Linear(checkpoint, f"{prefix}.attn.qkv_proj", dtype) self.out = _Linear(checkpoint, f"{prefix}.attn.out_proj", dtype) self.fc1 = _Linear(checkpoint, f"{prefix}.mlp.fc1", dtype) self.fc2 = _Linear(checkpoint, f"{prefix}.mlp.fc2", dtype) self.register_buffer("norm1", checkpoint.tensor(f"{prefix}.norm1.weight", dtype=dtype), persistent=False) self.register_buffer("norm2", checkpoint.tensor(f"{prefix}.norm2.weight", dtype=dtype), persistent=False) self.register_buffer("q_norm", checkpoint.tensor(f"{prefix}.attn.q_norm.weight", dtype=dtype), persistent=False) self.register_buffer("k_norm", checkpoint.tensor(f"{prefix}.attn.k_norm.weight", dtype=dtype), persistent=False) def forward(self, x: torch.Tensor) -> torch.Tensor: normalized = rms_norm(x, self.norm1, 1e-5) sequence = x.shape[0] q, k, v = self.qkv(normalized).split(7168, dim=-1) q = rms_norm(q.view(1, sequence, 56, 128), self.q_norm, 1e-5).transpose(1, 2) k = rms_norm(k.view(1, sequence, 56, 128), self.k_norm, 1e-5).transpose(1, 2) v = v.view(1, sequence, 56, 128).transpose(1, 2) x = x + self.out(functional.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(sequence, 7168)) gate, up = self.fc1(rms_norm(x, self.norm2, 1e-5)).chunk(2, dim=-1) return x + self.fc2(functional.silu(gate) * up) class H3TokenRefiner(nn.Module): """Project Qwen layer-50 states and refine them for H3 T2V.""" def __init__(self, checkpoint: H3Checkpoint, dtype: torch.dtype = torch.bfloat16): super().__init__() self.register_buffer("condition_weight", checkpoint.tensor("condition_proj.weight", dtype=dtype), persistent=False) self.register_buffer("condition_bias", checkpoint.tensor("condition_proj.bias", dtype=dtype), persistent=False) self.blocks = nn.ModuleList(_RefinerBlock(checkpoint, f"token_refiner.blocks.{index}", dtype) for index in range(2)) self.register_buffer("final_norm", checkpoint.tensor("token_refiner.final_norm.weight", dtype=dtype), persistent=False) @torch.inference_mode() def forward(self, qwen_states: torch.Tensor) -> torch.Tensor: x = functional.linear(qwen_states[0].to(self.condition_weight.dtype), self.condition_weight, self.condition_bias) for block in self.blocks: x = block(x) return rms_norm(x, self.final_norm, 1e-5).unsqueeze(0)