返回 JoyAI-Echo
model.py
1 from enum import Enum
2
3 import torch
4
5 from ltx_core.guidance.perturbations import BatchedPerturbationConfig
6 from ltx_core.model.transformer.adaln import AdaLayerNormSingle, adaln_embedding_coefficient
7 from ltx_core.model.transformer.attention import AttentionCallable, AttentionFunction
8 from ltx_core.model.transformer.modality import Modality
9 from ltx_core.model.transformer.rope import LTXRopeType
10 from ltx_core.model.transformer.transformer import BasicAVTransformerBlock, TransformerConfig
11 from ltx_core.model.transformer.transformer_args import (
12 MultiModalTransformerArgsPreprocessor,
13 TransformerArgs,
14 TransformerArgsPreprocessor,
15 )
16 from ltx_core.utils import to_denoised
17
18
19 class LTXModelType(Enum):
20 AudioVideo = "ltx av model"
21 VideoOnly = "ltx video only model"
22 AudioOnly = "ltx audio only model"
23
24 def is_video_enabled(self) -> bool:
25 return self in (LTXModelType.AudioVideo, LTXModelType.VideoOnly)
26
27 def is_audio_enabled(self) -> bool:
28 return self in (LTXModelType.AudioVideo, LTXModelType.AudioOnly)
29
30
31 class LTXModel(torch.nn.Module):
32 """
33 LTX model transformer implementation.
34 This class implements the transformer blocks for the LTX model.
35 """
36
37 def __init__( # noqa: PLR0913
38 self,
39 *,
40 model_type: LTXModelType = LTXModelType.AudioVideo,
41 num_attention_heads: int = 32,
42 attention_head_dim: int = 128,
43 in_channels: int = 128,
44 out_channels: int = 128,
45 num_layers: int = 48,
46 cross_attention_dim: int = 4096,
47 norm_eps: float = 1e-06,
48 attention_type: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
49 positional_embedding_theta: float = 10000.0,
50 positional_embedding_max_pos: list[int] | None = None,
51 timestep_scale_multiplier: int = 1000,
52 use_middle_indices_grid: bool = True,
53 audio_num_attention_heads: int = 32,
54 audio_attention_head_dim: int = 64,
55 audio_in_channels: int = 128,
56 audio_out_channels: int = 128,
57 audio_cross_attention_dim: int = 2048,
58 audio_positional_embedding_max_pos: list[int] | None = None,
59 av_ca_timestep_scale_multiplier: int = 1,
60 rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
61 double_precision_rope: bool = False,
62 apply_gated_attention: bool = False,
63 caption_projection: torch.nn.Module | None = None,
64 audio_caption_projection: torch.nn.Module | None = None,
65 cross_attention_adaln: bool = False,
66 ):
67 super().__init__()
68 self._enable_gradient_checkpointing = False
69 self.cross_attention_adaln = cross_attention_adaln
70 self.use_middle_indices_grid = use_middle_indices_grid
71 self.rope_type = rope_type
72 self.double_precision_rope = double_precision_rope
73 self.timestep_scale_multiplier = timestep_scale_multiplier
74 self.positional_embedding_theta = positional_embedding_theta
75 self.model_type = model_type
76 cross_pe_max_pos = None
77 if model_type.is_video_enabled():
78 if positional_embedding_max_pos is None:
79 positional_embedding_max_pos = [20, 2048, 2048]
80 self.positional_embedding_max_pos = positional_embedding_max_pos
81 self.num_attention_heads = num_attention_heads
82 self.inner_dim = num_attention_heads * attention_head_dim
83 self._init_video(
84 in_channels=in_channels,
85 out_channels=out_channels,
86 norm_eps=norm_eps,
87 caption_projection=caption_projection,
88 )
89
90 if model_type.is_audio_enabled():
91 if audio_positional_embedding_max_pos is None:
92 audio_positional_embedding_max_pos = [20]
93 self.audio_positional_embedding_max_pos = audio_positional_embedding_max_pos
94 self.audio_num_attention_heads = audio_num_attention_heads
95 self.audio_inner_dim = self.audio_num_attention_heads * audio_attention_head_dim
96 self._init_audio(
97 in_channels=audio_in_channels,
98 out_channels=audio_out_channels,
99 norm_eps=norm_eps,
100 caption_projection=audio_caption_projection,
101 )
102
103 if model_type.is_video_enabled() and model_type.is_audio_enabled():
104 cross_pe_max_pos = max(self.positional_embedding_max_pos[0], self.audio_positional_embedding_max_pos[0])
105 self.av_ca_timestep_scale_multiplier = av_ca_timestep_scale_multiplier
106 self.audio_cross_attention_dim = audio_cross_attention_dim
107 self._init_audio_video(num_scale_shift_values=4)
108
109 self._init_preprocessors(cross_pe_max_pos)
110 # Initialize transformer blocks
111 self._init_transformer_blocks(
112 num_layers=num_layers,
113 attention_head_dim=attention_head_dim if model_type.is_video_enabled() else 0,
114 cross_attention_dim=cross_attention_dim,
115 audio_attention_head_dim=audio_attention_head_dim if model_type.is_audio_enabled() else 0,
116 audio_cross_attention_dim=audio_cross_attention_dim,
117 norm_eps=norm_eps,
118 attention_type=attention_type,
119 apply_gated_attention=apply_gated_attention,
120 )
121
122 @property
123 def _adaln_embedding_coefficient(self) -> int:
124 return adaln_embedding_coefficient(self.cross_attention_adaln)
125
126 def _init_video(
127 self,
128 in_channels: int,
129 out_channels: int,
130 norm_eps: float,
131 caption_projection: torch.nn.Module | None = None,
132 ) -> None:
133 """Initialize video-specific components."""
134 # Video input components
135 self.patchify_proj = torch.nn.Linear(in_channels, self.inner_dim, bias=True)
136 if caption_projection is not None:
137 self.caption_projection = caption_projection
138
139 self.adaln_single = AdaLayerNormSingle(self.inner_dim, embedding_coefficient=self._adaln_embedding_coefficient)
140
141 self.prompt_adaln_single = (
142 AdaLayerNormSingle(self.inner_dim, embedding_coefficient=2) if self.cross_attention_adaln else None
143 )
144
145 # Video output components
146 self.scale_shift_table = torch.nn.Parameter(torch.empty(2, self.inner_dim))
147 self.norm_out = torch.nn.LayerNorm(self.inner_dim, elementwise_affine=False, eps=norm_eps)
148 self.proj_out = torch.nn.Linear(self.inner_dim, out_channels)
149
150 def _init_audio(
151 self,
152 in_channels: int,
153 out_channels: int,
154 norm_eps: float,
155 caption_projection: torch.nn.Module | None = None,
156 ) -> None:
157 """Initialize audio-specific components."""
158
159 # Audio input components
160 self.audio_patchify_proj = torch.nn.Linear(in_channels, self.audio_inner_dim, bias=True)
161 if caption_projection is not None:
162 self.audio_caption_projection = caption_projection
163
164 self.audio_adaln_single = AdaLayerNormSingle(
165 self.audio_inner_dim,
166 embedding_coefficient=self._adaln_embedding_coefficient,
167 )
168
169 self.audio_prompt_adaln_single = (
170 AdaLayerNormSingle(self.audio_inner_dim, embedding_coefficient=2) if self.cross_attention_adaln else None
171 )
172
173 # Audio output components
174 self.audio_scale_shift_table = torch.nn.Parameter(torch.empty(2, self.audio_inner_dim))
175 self.audio_norm_out = torch.nn.LayerNorm(self.audio_inner_dim, elementwise_affine=False, eps=norm_eps)
176 self.audio_proj_out = torch.nn.Linear(self.audio_inner_dim, out_channels)
177
178 def _init_audio_video(
179 self,
180 num_scale_shift_values: int,
181 ) -> None:
182 """Initialize audio-video cross-attention components."""
183 self.av_ca_video_scale_shift_adaln_single = AdaLayerNormSingle(
184 self.inner_dim,
185 embedding_coefficient=num_scale_shift_values,
186 )
187
188 self.av_ca_audio_scale_shift_adaln_single = AdaLayerNormSingle(
189 self.audio_inner_dim,
190 embedding_coefficient=num_scale_shift_values,
191 )
192
193 self.av_ca_a2v_gate_adaln_single = AdaLayerNormSingle(
194 self.inner_dim,
195 embedding_coefficient=1,
196 )
197
198 self.av_ca_v2a_gate_adaln_single = AdaLayerNormSingle(
199 self.audio_inner_dim,
200 embedding_coefficient=1,
201 )
202
203 def _init_preprocessors(
204 self,
205 cross_pe_max_pos: int | None = None,
206 ) -> None:
207 """Initialize preprocessors for LTX."""
208
209 if self.model_type.is_video_enabled() and self.model_type.is_audio_enabled():
210 self.video_args_preprocessor = MultiModalTransformerArgsPreprocessor(
211 patchify_proj=self.patchify_proj,
212 adaln=self.adaln_single,
213 cross_scale_shift_adaln=self.av_ca_video_scale_shift_adaln_single,
214 cross_gate_adaln=self.av_ca_a2v_gate_adaln_single,
215 inner_dim=self.inner_dim,
216 max_pos=self.positional_embedding_max_pos,
217 num_attention_heads=self.num_attention_heads,
218 cross_pe_max_pos=cross_pe_max_pos,
219 use_middle_indices_grid=self.use_middle_indices_grid,
220 audio_cross_attention_dim=self.audio_cross_attention_dim,
221 timestep_scale_multiplier=self.timestep_scale_multiplier,
222 double_precision_rope=self.double_precision_rope,
223 positional_embedding_theta=self.positional_embedding_theta,
224 rope_type=self.rope_type,
225 av_ca_timestep_scale_multiplier=self.av_ca_timestep_scale_multiplier,
226 caption_projection=getattr(self, "caption_projection", None),
227 prompt_adaln=getattr(self, "prompt_adaln_single", None),
228 )
229 self.audio_args_preprocessor = MultiModalTransformerArgsPreprocessor(
230 patchify_proj=self.audio_patchify_proj,
231 adaln=self.audio_adaln_single,
232 cross_scale_shift_adaln=self.av_ca_audio_scale_shift_adaln_single,
233 cross_gate_adaln=self.av_ca_v2a_gate_adaln_single,
234 inner_dim=self.audio_inner_dim,
235 max_pos=self.audio_positional_embedding_max_pos,
236 num_attention_heads=self.audio_num_attention_heads,
237 cross_pe_max_pos=cross_pe_max_pos,
238 use_middle_indices_grid=self.use_middle_indices_grid,
239 audio_cross_attention_dim=self.audio_cross_attention_dim,
240 timestep_scale_multiplier=self.timestep_scale_multiplier,
241 double_precision_rope=self.double_precision_rope,
242 positional_embedding_theta=self.positional_embedding_theta,
243 rope_type=self.rope_type,
244 av_ca_timestep_scale_multiplier=self.av_ca_timestep_scale_multiplier,
245 caption_projection=getattr(self, "audio_caption_projection", None),
246 prompt_adaln=getattr(self, "audio_prompt_adaln_single", None),
247 )
248 elif self.model_type.is_video_enabled():
249 self.video_args_preprocessor = TransformerArgsPreprocessor(
250 patchify_proj=self.patchify_proj,
251 adaln=self.adaln_single,
252 inner_dim=self.inner_dim,
253 max_pos=self.positional_embedding_max_pos,
254 num_attention_heads=self.num_attention_heads,
255 use_middle_indices_grid=self.use_middle_indices_grid,
256 timestep_scale_multiplier=self.timestep_scale_multiplier,
257 double_precision_rope=self.double_precision_rope,
258 positional_embedding_theta=self.positional_embedding_theta,
259 rope_type=self.rope_type,
260 caption_projection=getattr(self, "caption_projection", None),
261 prompt_adaln=getattr(self, "prompt_adaln_single", None),
262 )
263 elif self.model_type.is_audio_enabled():
264 self.audio_args_preprocessor = TransformerArgsPreprocessor(
265 patchify_proj=self.audio_patchify_proj,
266 adaln=self.audio_adaln_single,
267 inner_dim=self.audio_inner_dim,
268 max_pos=self.audio_positional_embedding_max_pos,
269 num_attention_heads=self.audio_num_attention_heads,
270 use_middle_indices_grid=self.use_middle_indices_grid,
271 timestep_scale_multiplier=self.timestep_scale_multiplier,
272 double_precision_rope=self.double_precision_rope,
273 positional_embedding_theta=self.positional_embedding_theta,
274 rope_type=self.rope_type,
275 caption_projection=getattr(self, "audio_caption_projection", None),
276 prompt_adaln=getattr(self, "audio_prompt_adaln_single", None),
277 )
278
279 def _init_transformer_blocks(
280 self,
281 num_layers: int,
282 attention_head_dim: int,
283 cross_attention_dim: int,
284 audio_attention_head_dim: int,
285 audio_cross_attention_dim: int,
286 norm_eps: float,
287 attention_type: AttentionFunction | AttentionCallable,
288 apply_gated_attention: bool,
289 ) -> None:
290 """Initialize transformer blocks for LTX."""
291 video_config = (
292 TransformerConfig(
293 dim=self.inner_dim,
294 heads=self.num_attention_heads,
295 d_head=attention_head_dim,
296 context_dim=cross_attention_dim,
297 apply_gated_attention=apply_gated_attention,
298 cross_attention_adaln=self.cross_attention_adaln,
299 )
300 if self.model_type.is_video_enabled()
301 else None
302 )
303 audio_config = (
304 TransformerConfig(
305 dim=self.audio_inner_dim,
306 heads=self.audio_num_attention_heads,
307 d_head=audio_attention_head_dim,
308 context_dim=audio_cross_attention_dim,
309 apply_gated_attention=apply_gated_attention,
310 cross_attention_adaln=self.cross_attention_adaln,
311 )
312 if self.model_type.is_audio_enabled()
313 else None
314 )
315 self.transformer_blocks = torch.nn.ModuleList(
316 [
317 BasicAVTransformerBlock(
318 idx=idx,
319 num_layers=num_layers,
320 video=video_config,
321 audio=audio_config,
322 rope_type=self.rope_type,
323 norm_eps=norm_eps,
324 attention_function=attention_type,
325 )
326 for idx in range(num_layers)
327 ]
328 )
329
330 def enable_action_conditioning(self, action_config) -> None:
331 """Attach the optional pure-UCPE branch after base weights are loaded."""
332 from ltx_core.model.transformer.transformer import ActionBlockConfig
333
334 if not isinstance(action_config, ActionBlockConfig):
335 raise TypeError(f"action_config must be ActionBlockConfig, got {type(action_config)}")
336 if not self.model_type.is_video_enabled():
337 raise RuntimeError("UCPE conditioning requires a video-enabled model")
338 video_config = TransformerConfig(
339 dim=self.inner_dim,
340 heads=self.num_attention_heads,
341 d_head=self.inner_dim // self.num_attention_heads,
342 context_dim=0,
343 )
344 for block in self.transformer_blocks:
345 block._init_action_params(video_config, action_config)
346
347 def set_gradient_checkpointing(self, enable: bool) -> None:
348 """Enable or disable gradient checkpointing for transformer blocks.
349 Gradient checkpointing trades compute for memory by recomputing activations
350 during the backward pass instead of storing them. This can significantly
351 reduce memory usage at the cost of ~20-30% slower training.
352 Args:
353 enable: Whether to enable gradient checkpointing
354 """
355 self._enable_gradient_checkpointing = enable
356
357 def _process_transformer_blocks(
358 self,
359 video: TransformerArgs | None,
360 audio: TransformerArgs | None,
361 perturbations: BatchedPerturbationConfig,
362 action_cond: dict | None = None,
363 kv_caches: list[dict] | None = None,
364 current_video_token_start: int = 0,
365 current_audio_token_start: int = 0,
366 ) -> tuple[TransformerArgs, TransformerArgs]:
367 """Process transformer blocks for LTXAV."""
368
369 ucpe_viewmats = action_cond.get("ucpe_viewmats") if action_cond else None
370 ucpe_Ks = action_cond.get("ucpe_Ks") if action_cond else None
371 for layer_index, block in enumerate(self.transformer_blocks):
372 layer_cache = kv_caches[layer_index] if kv_caches is not None else None
373 if self._enable_gradient_checkpointing and self.training:
374 # Use gradient checkpointing to save memory during training.
375 # With use_reentrant=False, we can pass dataclasses directly -
376 # PyTorch will track all tensor leaves in the computation graph.
377 video, audio = torch.utils.checkpoint.checkpoint(
378 block,
379 video,
380 audio,
381 perturbations,
382 ucpe_viewmats,
383 ucpe_Ks,
384 layer_cache,
385 current_video_token_start,
386 current_audio_token_start,
387 use_reentrant=False,
388 )
389 else:
390 video, audio = block(
391 video=video,
392 audio=audio,
393 perturbations=perturbations,
394 ucpe_viewmats=ucpe_viewmats,
395 ucpe_Ks=ucpe_Ks,
396 kv_cache=layer_cache,
397 current_video_token_start=current_video_token_start,
398 current_audio_token_start=current_audio_token_start,
399 )
400
401 return video, audio
402
403 def _process_output(
404 self,
405 scale_shift_table: torch.Tensor,
406 norm_out: torch.nn.LayerNorm,
407 proj_out: torch.nn.Linear,
408 x: torch.Tensor,
409 embedded_timestep: torch.Tensor,
410 ) -> torch.Tensor:
411 """Process output for LTXV."""
412 # Apply scale-shift modulation
413 scale_shift_values = (
414 scale_shift_table[None, None].to(device=x.device, dtype=x.dtype) + embedded_timestep[:, :, None]
415 )
416 shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1]
417
418 x = norm_out(x)
419 x = x * (1 + scale) + shift
420 x = proj_out(x)
421 return x
422
423 def forward(
424 self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig,
425 action_cond: dict | None = None,
426 kv_caches: list[dict] | None = None,
427 current_video_token_start: int = 0,
428 current_audio_token_start: int = 0,
429 ) -> tuple[torch.Tensor, torch.Tensor]:
430 """
431 Forward pass for LTX models.
432 Returns:
433 Processed output tensors
434 """
435 if not self.model_type.is_video_enabled() and video is not None:
436 raise ValueError("Video is not enabled for this model")
437 if not self.model_type.is_audio_enabled() and audio is not None:
438 raise ValueError("Audio is not enabled for this model")
439
440 video_args = self.video_args_preprocessor.prepare(video, audio) if video is not None else None
441 audio_args = self.audio_args_preprocessor.prepare(audio, video) if audio is not None else None
442 # Process transformer blocks
443 video_out, audio_out = self._process_transformer_blocks(
444 video=video_args,
445 audio=audio_args,
446 perturbations=perturbations,
447 action_cond=action_cond,
448 kv_caches=kv_caches,
449 current_video_token_start=current_video_token_start,
450 current_audio_token_start=current_audio_token_start,
451 )
452
453 # Process output
454 vx = (
455 self._process_output(
456 self.scale_shift_table, self.norm_out, self.proj_out, video_out.x, video_out.embedded_timestep
457 )
458 if video_out is not None
459 else None
460 )
461 ax = (
462 self._process_output(
463 self.audio_scale_shift_table,
464 self.audio_norm_out,
465 self.audio_proj_out,
466 audio_out.x,
467 audio_out.embedded_timestep,
468 )
469 if audio_out is not None
470 else None
471 )
472 return vx, ax
473
474 def init_av_kv_caches( # noqa: PLR0913
475 self,
476 batch_size: int,
477 max_video_tokens: int,
478 max_audio_tokens: int,
479 text_seq_len: int,
480 device: torch.device,
481 dtype: torch.dtype,
482 video_local_attn_tokens: int = -1,
483 video_sink_tokens: int = 0,
484 video_ucpe_local_attn_tokens: int | None = None,
485 video_ucpe_sink_tokens: int | None = None,
486 audio_local_attn_tokens: int = -1,
487 audio_sink_tokens: int = 0,
488 ) -> list[dict]:
489 """Allocate inference-only per-layer AV caches."""
490
491 def allocate(max_tokens: int, dim: int, local: int, sink: int) -> dict:
492 capacity = max_tokens if local < 0 else min(max_tokens, local)
493 if capacity <= 0:
494 raise ValueError("KV cache capacity must be positive")
495 return {
496 "k": torch.zeros(batch_size, capacity, dim, device=device, dtype=dtype),
497 "v": torch.zeros(batch_size, capacity, dim, device=device, dtype=dtype),
498 "positions": torch.full((capacity,), -1, device=device, dtype=torch.long),
499 "length": 0,
500 "local_attn_size": local,
501 "sink_tokens": sink,
502 }
503
504 def allocate_text(dim: int) -> dict:
505 return {
506 "k": torch.zeros(batch_size, text_seq_len, dim, device=device, dtype=dtype),
507 "v": torch.zeros(batch_size, text_seq_len, dim, device=device, dtype=dtype),
508 "length": 0,
509 "is_init": False,
510 }
511
512 ucpe_local = video_local_attn_tokens if video_ucpe_local_attn_tokens is None else video_ucpe_local_attn_tokens
513 ucpe_sink = video_sink_tokens if video_ucpe_sink_tokens is None else video_ucpe_sink_tokens
514 caches: list[dict] = []
515 for block in self.transformer_blocks:
516 layer: dict = {}
517 if self.model_type.is_video_enabled():
518 layer["video_self"] = allocate(max_video_tokens, self.inner_dim, video_local_attn_tokens, video_sink_tokens)
519 layer["video_text"] = allocate_text(self.inner_dim)
520 if block.action_ucpe_enabled:
521 layer["video_ucpe"] = allocate(
522 max_video_tokens,
523 block.ucpe_num_heads * block.ucpe_head_dim,
524 ucpe_local,
525 ucpe_sink,
526 )
527 if self.model_type.is_audio_enabled() and max_audio_tokens > 0:
528 layer["audio_self"] = allocate(max_audio_tokens, self.audio_inner_dim, audio_local_attn_tokens, audio_sink_tokens)
529 layer["audio_text"] = allocate_text(self.audio_inner_dim)
530 if self.model_type.is_video_enabled() and self.model_type.is_audio_enabled() and max_audio_tokens > 0:
531 layer["a2v"] = allocate(max_audio_tokens, self.audio_inner_dim, audio_local_attn_tokens, audio_sink_tokens)
532 layer["v2a"] = allocate(max_video_tokens, self.audio_inner_dim, video_local_attn_tokens, video_sink_tokens)
533 caches.append(layer)
534 return caches
535
536
537 class LegacyX0Model(torch.nn.Module):
538 """
539 Legacy X0 model implementation.
540 Returns fully denoised output based on the velocities produced by the base model.
541 """
542
543 def __init__(self, velocity_model: LTXModel):
544 super().__init__()
545 self.velocity_model = velocity_model
546
547 def forward(
548 self,
549 video: Modality | None,
550 audio: Modality | None,
551 perturbations: BatchedPerturbationConfig,
552 sigma: float,
553 action_cond: dict | None = None,
554 ) -> tuple[torch.Tensor | None, torch.Tensor | None]:
555 """
556 Denoise the video and audio according to the sigma.
557 Returns:
558 Denoised video and audio
559 """
560 vx, ax = self.velocity_model(video, audio, perturbations, action_cond=action_cond)
561 denoised_video = to_denoised(video.latent, vx, sigma) if vx is not None else None
562 denoised_audio = to_denoised(audio.latent, ax, sigma) if ax is not None else None
563 return denoised_video, denoised_audio
564
565
566 class X0Model(torch.nn.Module):
567 """
568 X0 model implementation.
569 Returns fully denoised outputs based on the velocities produced by the base model.
570 Applies scaled denoising to the video and audio according to the timesteps = sigma * denoising_mask.
571 """
572
573 def __init__(self, velocity_model: LTXModel):
574 super().__init__()
575 self.velocity_model = velocity_model
576
577 def forward(
578 self,
579 video: Modality | None,
580 audio: Modality | None,
581 perturbations: BatchedPerturbationConfig,
582 action_cond: dict | None = None,
583 ) -> tuple[torch.Tensor | None, torch.Tensor | None]:
584 """
585 Denoise the video and audio according to the sigma.
586 Returns:
587 Denoised video and audio
588 """
589 vx, ax = self.velocity_model(video, audio, perturbations, action_cond=action_cond)
590 denoised_video = to_denoised(video.latent, vx, video.timesteps) if vx is not None else None
591 denoised_audio = to_denoised(audio.latent, ax, audio.timesteps) if ax is not None else None
592 return denoised_video, denoised_audio
593
593 lines PYTHON