"""Final H3 curve-AdaLN and video/audio output heads.""" import torch import torch.nn.functional as functional from torch import nn from .attention import rms_norm from .checkpoint import H3Checkpoint class H3FinalLayer(nn.Module): def __init__( self, curve_table: torch.Tensor, norm_weight: torch.Tensor, adaln_weight: torch.Tensor, adaln_bias: torch.Tensor, video_weight: torch.Tensor, video_bias: torch.Tensor, audio_weight: torch.Tensor, audio_bias: torch.Tensor, *, hidden_size: int = 5376, eps: float = 1e-5, ): super().__init__() if adaln_weight.shape != (2 * hidden_size, curve_table.shape[1]): raise ValueError("Unexpected final H3 AdaLN projection dimensions.") self.hidden_size = hidden_size self.eps = eps self.register_buffer("curve_table", curve_table.to(torch.float32), persistent=False) self.register_buffer("norm_weight", norm_weight, persistent=False) self.register_buffer("adaln_weight", adaln_weight.to(torch.float32), persistent=False) self.register_buffer("adaln_bias", adaln_bias.to(torch.float32), persistent=False) self.register_buffer("video_weight", video_weight.to(torch.float32), persistent=False) self.register_buffer("video_bias", video_bias.to(torch.float32), persistent=False) self.register_buffer("audio_weight", audio_weight.to(torch.float32), persistent=False) self.register_buffer("audio_bias", audio_bias.to(torch.float32), persistent=False) @classmethod def from_checkpoint(cls, checkpoint: H3Checkpoint, *, output_dtype=torch.bfloat16): return cls( checkpoint.tensor("adaln_t_table", dtype=torch.float32), checkpoint.tensor("final_layer.norm.weight", dtype=output_dtype), checkpoint.tensor("final_layer.adaln_proj.linear.weight", dtype=torch.float32), checkpoint.tensor("final_layer.adaln_proj.linear.bias", dtype=torch.float32), checkpoint.tensor("final_layer.video_out.weight", dtype=torch.float32), checkpoint.tensor("final_layer.video_out.bias", dtype=torch.float32), checkpoint.tensor("final_layer.audio_out.weight", dtype=torch.float32), checkpoint.tensor("final_layer.audio_out.bias", dtype=torch.float32), ) def forward( self, hidden: torch.Tensor, timesteps: torch.Tensor, video_segment: tuple[int, int, int], audio_segment: tuple[int, int, int], ) -> tuple[torch.Tensor, torch.Tensor]: position = timesteps.float().clamp(0, 1) * (self.curve_table.shape[0] - 1) lower = position.floor().long().clamp(max=self.curve_table.shape[0] - 2) embedding = torch.lerp(self.curve_table[lower], self.curve_table[lower + 1], (position - lower).unsqueeze(1)) shift, scale = functional.linear(embedding, self.adaln_weight, self.adaln_bias).chunk(2, dim=-1) video_start, video_stop, video_row = video_segment audio_start, audio_stop, audio_row = audio_segment normalized = rms_norm(hidden, self.norm_weight, self.eps) video_hidden = normalized[video_start:video_stop] * (1 + scale[video_row].to(hidden.dtype)) + shift[video_row].to(hidden.dtype) audio_hidden = normalized[audio_start:audio_stop] * (1 + scale[audio_row].to(hidden.dtype)) + shift[audio_row].to(hidden.dtype) return ( functional.linear(video_hidden.float(), self.video_weight, self.video_bias), functional.linear(audio_hidden.float(), self.audio_weight, self.audio_bias), )