返回 JoyAI-Echo
feature_extractor.py
根目录 / ltx-core / src / ltx_core / text_encoders / gemma / feature_extractor.py
1 import math
2
3 import torch
4 from einops import rearrange
5 from torch import nn
6
7 # ---------------------------------------------------------------------------
8 # Normalization functions
9 # ---------------------------------------------------------------------------
10
11
12 def _norm_and_concat_padded_batch(
13 encoded_text: torch.Tensor,
14 sequence_lengths: torch.Tensor,
15 padding_side: str = "right",
16 ) -> torch.Tensor:
17 """Normalize and flatten multi-layer hidden states, respecting padding.
18 Performs per-batch, per-layer normalization using masked mean and range,
19 then concatenates across the layer dimension.
20 Args:
21 encoded_text: Hidden states of shape [batch, seq_len, hidden_dim, num_layers].
22 sequence_lengths: Number of valid (non-padded) tokens per batch item.
23 padding_side: Whether padding is on "left" or "right".
24 Returns:
25 Normalized tensor of shape [batch, seq_len, hidden_dim * num_layers],
26 with padded positions zeroed out.
27 """
28 b, t, d, l = encoded_text.shape # noqa: E741
29 device = encoded_text.device
30
31 token_indices = torch.arange(t, device=device)[None, :]
32
33 if padding_side == "right":
34 mask = token_indices < sequence_lengths[:, None]
35 elif padding_side == "left":
36 start_indices = t - sequence_lengths[:, None]
37 mask = token_indices >= start_indices
38 else:
39 raise ValueError(f"padding_side must be 'left' or 'right', got {padding_side}")
40
41 mask = rearrange(mask, "b t -> b t 1 1")
42
43 eps = 1e-6
44
45 masked = encoded_text.masked_fill(~mask, 0.0)
46 denom = (sequence_lengths * d).view(b, 1, 1, 1)
47 mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps)
48
49 x_min = encoded_text.masked_fill(~mask, float("inf")).amin(dim=(1, 2), keepdim=True)
50 x_max = encoded_text.masked_fill(~mask, float("-inf")).amax(dim=(1, 2), keepdim=True)
51 range_ = x_max - x_min
52
53 normed = 8 * (encoded_text - mean) / (range_ + eps)
54 normed = normed.reshape(b, t, -1)
55
56 mask_flattened = rearrange(mask, "b t 1 1 -> b t 1").expand(-1, -1, d * l)
57 normed = normed.masked_fill(~mask_flattened, 0.0)
58
59 return normed
60
61
62 def norm_and_concat_per_token_rms(
63 encoded_text: torch.Tensor,
64 attention_mask: torch.Tensor,
65 ) -> torch.Tensor:
66 """Per-token RMSNorm normalization for V2 models.
67 Args:
68 encoded_text: [B, T, D, L]
69 attention_mask: [B, T] binary mask
70 Returns:
71 [B, T, D*L] normalized tensor with padding zeroed out.
72 """
73 B, T, D, L = encoded_text.shape # noqa: N806
74 variance = torch.mean(encoded_text**2, dim=2, keepdim=True) # [B,T,1,L]
75 normed = encoded_text * torch.rsqrt(variance + 1e-6)
76 normed = normed.reshape(B, T, D * L)
77 mask_3d = attention_mask.bool().unsqueeze(-1) # [B, T, 1]
78 return torch.where(mask_3d, normed, torch.zeros_like(normed))
79
80
81 def _rescale_norm(x: torch.Tensor, target_dim: int, source_dim: int) -> torch.Tensor:
82 """Rescale normalization: x * sqrt(target_dim / source_dim)."""
83 return x * math.sqrt(target_dim / source_dim)
84
85
86 # ---------------------------------------------------------------------------
87 # Feature extractor variants
88 # ---------------------------------------------------------------------------
89
90
91 class FeatureExtractorV1(nn.Module):
92 """19B: per-segment norm -> aggregate_embed -> 3840"""
93
94 def __init__(self, aggregate_embed: nn.Module, is_av: bool = False):
95 super().__init__()
96 self.aggregate_embed = aggregate_embed
97 self.is_av = is_av
98
99 def forward(
100 self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, padding_side: str = "left"
101 ) -> tuple[torch.Tensor, torch.Tensor | None]:
102 encoded = torch.stack(hidden_states, dim=-1) if isinstance(hidden_states, (list, tuple)) else hidden_states
103 dtype = encoded.dtype
104 sequence_lengths = attention_mask.sum(dim=-1)
105 normed = _norm_and_concat_padded_batch(encoded, sequence_lengths, padding_side)
106 features = self.aggregate_embed(normed.to(dtype))
107 if self.is_av:
108 return features, features
109 return features, None
110
111
112 class FeatureExtractorV2(nn.Module):
113 """22B: per-token RMS norm → rescale → dual aggregate embeds"""
114
115 def __init__(
116 self,
117 video_aggregate_embed: nn.Linear,
118 embedding_dim: int,
119 audio_aggregate_embed: nn.Linear | None = None,
120 ):
121 super().__init__()
122 self.video_aggregate_embed = video_aggregate_embed
123 self.audio_aggregate_embed = audio_aggregate_embed
124 self.embedding_dim = embedding_dim
125
126 def forward(
127 self,
128 hidden_states: torch.Tensor,
129 attention_mask: torch.Tensor,
130 padding_side: str = "left", # noqa: ARG002
131 ) -> tuple[torch.Tensor, torch.Tensor | None]:
132 encoded = torch.stack(hidden_states, dim=-1) if isinstance(hidden_states, (list, tuple)) else hidden_states
133 normed = norm_and_concat_per_token_rms(encoded, attention_mask)
134 normed = normed.to(encoded.dtype)
135 v_dim = self.video_aggregate_embed.out_features
136 video = self.video_aggregate_embed(_rescale_norm(normed, v_dim, self.embedding_dim))
137 audio = None
138 if self.audio_aggregate_embed is not None:
139 a_dim = self.audio_aggregate_embed.out_features
140 audio = self.audio_aggregate_embed(_rescale_norm(normed, a_dim, self.embedding_dim))
141 return video, audio
142
142 lines PYTHON