33
44from __future__ import annotations
55
6+ import gc
7+ from dataclasses import dataclass
8+ from typing import Any
9+
10+ import torch
11+
12+ from fastvideo .configs .models .vaes .minimax_h3_audio import MiniMaxH3AudioVAEArchConfig
13+ from fastvideo .configs .models .vaes .minimax_h3_video import MiniMaxH3VideoVAEArchConfig
614from fastvideo .configs .pipelines .minimax_h3 import MiniMaxH3PipelineConfig
715from fastvideo .fastvideo_args import FastVideoArgs
16+ from fastvideo .logger import init_logger
817from fastvideo .pipelines .basic .minimax_h3 .stages import (
918 MiniMaxH3AudioDecodingStage ,
1019 MiniMaxH3ConditioningStage ,
1524)
1625from fastvideo .pipelines .composed_pipeline_base import ComposedPipelineBase
1726from fastvideo .pipelines .lora_pipeline import LoRAPipeline
27+ from fastvideo .pipelines .pipeline_batch_info import ForwardBatch
28+
29+ logger = init_logger (__name__ )
30+
31+ # Same split as the MLX runtime: condition, release the ~66 GB Qwen3-VL stack,
32+ # then load DiT + VAEs. Keeping them resident together OOMs unified-memory
33+ # boxes (GB10 / Spark) even though host offload is correctly disabled there.
34+ _DENOISE_MODULE_NAMES = ("vae" , "audio_vae" , "transformer" )
35+
36+
37+ @dataclass (frozen = True )
38+ class _H3VideoGeometry :
39+ spatial_compression_ratio : int
40+ latent_channels : int
41+
42+
43+ @dataclass (frozen = True )
44+ class _H3AudioGeometry :
45+ sampling_rate : int
46+
47+
48+ def _default_video_geometry () -> _H3VideoGeometry :
49+ arch = MiniMaxH3VideoVAEArchConfig ()
50+ return _H3VideoGeometry (
51+ spatial_compression_ratio = int (arch .spatial_compression_ratio ),
52+ latent_channels = int (arch .latent_channels ),
53+ )
54+
55+
56+ def _default_audio_geometry () -> _H3AudioGeometry :
57+ return _H3AudioGeometry (sampling_rate = int (MiniMaxH3AudioVAEArchConfig ().sampling_rate ))
1858
1959
2060class MiniMaxH3BasePipeline (LoRAPipeline , ComposedPipelineBase ):
@@ -53,6 +93,11 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase):
5393 "audio_scheduler" ,
5494 ]
5595
96+ def __init__ (self , * args : Any , ** kwargs : Any ) -> None :
97+ self ._ref2va = False
98+ self ._denoise_stages_ready = False
99+ super ().__init__ (* args , ** kwargs )
100+
56101 @classmethod
57102 def get_hf_download_component_dirs (cls ) -> tuple [str , ...]:
58103 return tuple (sorted (cls ._extra_config_module_map .get (name , name ) for name in cls ._required_config_modules ))
@@ -67,18 +112,70 @@ def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
67112 if shift is None or float (shift ) != expected_shift :
68113 raise ValueError (f"MiniMax-H3 { modality } scheduler must expose shift={ expected_shift :g} , got { shift } ." )
69114
70- def _add_stages (self , * , ref2va : bool ) -> None :
71- transformer = self .get_module ("transformer" )
72- vae = self .get_module ("vae" )
73- audio_vae = self .get_module ("audio_vae" )
74- scheduler = self .get_module ("scheduler" )
75- audio_scheduler = self .get_module ("audio_scheduler" )
115+ def _defer_denoise_modules (self , fastvideo_args : FastVideoArgs ) -> bool :
116+ return bool (fastvideo_args .inference_mode ) and not bool (getattr (fastvideo_args , "training_mode" , False ))
117+
118+ def _denoise_modules_loaded (self ) -> bool :
119+ return all (self .get_module (name ) is not None for name in _DENOISE_MODULE_NAMES )
76120
121+ def load_modules (self ,
122+ fastvideo_args : FastVideoArgs ,
123+ loaded_modules : dict [str , torch .nn .Module ] | None = None ) -> dict [str , Any ]:
124+ """Load the Qwen3-VL conditioner first; defer DiT and VAEs until after encode."""
125+ if not self ._defer_denoise_modules (fastvideo_args ):
126+ return super ().load_modules (fastvideo_args , loaded_modules )
127+ if loaded_modules is not None and all (name in loaded_modules for name in _DENOISE_MODULE_NAMES ):
128+ return super ().load_modules (fastvideo_args , loaded_modules )
129+
130+ saved = list (self .required_config_modules )
131+ self ._required_config_modules = [name for name in saved if name not in _DENOISE_MODULE_NAMES ]
132+ try :
133+ logger .info ("Loading MiniMax-H3 condition modules first: %s" , self ._required_config_modules )
134+ return super ().load_modules (fastvideo_args , loaded_modules )
135+ finally :
136+ self ._required_config_modules = saved
137+
138+ def _load_denoise_modules (self , fastvideo_args : FastVideoArgs ) -> None :
139+ if self ._denoise_modules_loaded ():
140+ return
141+ saved = list (self .required_config_modules )
142+ self ._required_config_modules = [name for name in saved if name != "text_encoder" ]
143+ try :
144+ logger .info ("Loading MiniMax-H3 denoise modules after releasing the text encoder: %s" ,
145+ [name for name in self ._required_config_modules if name in _DENOISE_MODULE_NAMES ])
146+ loaded = super ().load_modules (fastvideo_args , loaded_modules = self .modules )
147+ for name , module in loaded .items ():
148+ self .add_module (name , module )
149+ finally :
150+ self ._required_config_modules = saved
151+
152+ def _release_text_encoder (self ) -> None :
153+ stage = self ._stage_name_mapping .get ("conditioning_stage" )
154+ if stage is not None :
155+ stage .conditioner = None
156+ encoder = self .modules .pop ("text_encoder" , None )
157+ if encoder is None :
158+ return
159+ logger .info ("Released MiniMax-H3 text encoder after conditioning" )
160+ del encoder
161+ gc .collect ()
162+ if torch .cuda .is_available ():
163+ torch .cuda .empty_cache ()
164+
165+ def _input_vae (self ) -> Any :
166+ return self .get_module ("vae" ) or _default_video_geometry ()
167+
168+ def _input_audio_vae (self , * , ref2va : bool ) -> Any | None :
169+ if not ref2va :
170+ return None
171+ return self .get_module ("audio_vae" ) or _default_audio_geometry ()
172+
173+ def _add_condition_stages (self , * , ref2va : bool ) -> None :
77174 self .add_stage (
78175 "input_preparation_stage" ,
79176 MiniMaxH3InputPreparationStage (
80- vae = vae ,
81- audio_vae = audio_vae if ref2va else None ,
177+ vae = self . _input_vae () ,
178+ audio_vae = self . _input_audio_vae ( ref2va = ref2va ) ,
82179 ref2va = ref2va ,
83180 ),
84181 )
@@ -91,6 +188,15 @@ def _add_stages(self, *, ref2va: bool) -> None:
91188 ref2va = ref2va ,
92189 ),
93190 )
191+
192+ def _add_denoise_stages (self , * , ref2va : bool ) -> None :
193+ transformer = self .get_module ("transformer" )
194+ vae = self .get_module ("vae" )
195+ audio_vae = self .get_module ("audio_vae" )
196+ scheduler = self .get_module ("scheduler" )
197+ audio_scheduler = self .get_module ("audio_scheduler" )
198+ if transformer is None or vae is None or audio_vae is None :
199+ raise RuntimeError ("MiniMax-H3 denoise stages require transformer, vae, and audio_vae to be loaded." )
94200 self .add_stage (
95201 "latent_preparation_stage" ,
96202 MiniMaxH3LatentPreparationStage (
@@ -111,6 +217,35 @@ def _add_stages(self, *, ref2va: bool) -> None:
111217 )
112218 self .add_stage ("video_decoding_stage" , MiniMaxH3VideoDecodingStage (vae = vae , transformer = transformer ))
113219 self .add_stage ("audio_decoding_stage" , MiniMaxH3AudioDecodingStage (audio_vae = audio_vae ))
220+ self ._denoise_stages_ready = True
221+
222+ def _add_stages (self , * , ref2va : bool ) -> None :
223+ self ._ref2va = ref2va
224+ self ._add_condition_stages (ref2va = ref2va )
225+ if self ._denoise_modules_loaded ():
226+ self ._add_denoise_stages (ref2va = ref2va )
227+
228+ def forward (self , batch : ForwardBatch , fastvideo_args : FastVideoArgs ) -> ForwardBatch :
229+ if not self .post_init_called :
230+ self .post_init ()
231+
232+ if self ._denoise_stages_ready :
233+ return super ().forward (batch , fastvideo_args )
234+
235+ logger .info ("Running MiniMax-H3 condition stages before loading DiT/VAE weights" )
236+ for stage in self .stages :
237+ batch = stage (batch , fastvideo_args )
238+ self ._release_text_encoder ()
239+ self ._load_denoise_modules (fastvideo_args )
240+ self ._add_denoise_stages (ref2va = self ._ref2va )
241+ for name in (
242+ "latent_preparation_stage" ,
243+ "denoising_stage" ,
244+ "video_decoding_stage" ,
245+ "audio_decoding_stage" ,
246+ ):
247+ batch = self ._stage_name_mapping [name ](batch , fastvideo_args )
248+ return batch
114249
115250
116251class MiniMaxH3Pipeline (MiniMaxH3BasePipeline ):
0 commit comments