@@ -377,10 +377,13 @@ def freeze(
377377 """
378378 modules = []
379379 if freeze_language_model and self .language_model :
380+ self .language_model .eval ()
380381 modules .append (self .language_model )
381382 if freeze_vision_model and self .vision_model is not None :
383+ self .vision_model .eval ()
382384 modules .append (self .vision_model )
383385 if freeze_vae_model and self .vae_model is not None :
386+ self .vae_model .eval ()
384387 modules .append (self .vae_model )
385388 if freeze_text_embed and self .language_model :
386389 modules .append (self .language_model .embedding .word_embeddings )
@@ -666,6 +669,10 @@ def forward(
666669 vit_token_seqlens : Optional [Tensor ] = None ,
667670 # VAE generation inputs
668671 padded_images : Optional [Tensor ] = None ,
672+ # padded_latent: Optional[Tensor] = None,
673+ # packed_latent_clean: Optional[Tensor] = None,
674+ # packed_latent_noise: Optional[Tensor] = None,
675+ # packed_timesteps_effective: Optional[Tensor] = None,
669676 patchified_vae_latent_shapes : Optional [List ] = None ,
670677 packed_latent_position_ids : Optional [Tensor ] = None ,
671678 packed_vae_token_indexes : Optional [Tensor ] = None ,
@@ -702,16 +709,18 @@ def forward(
702709 f"but got { tuple (packed_text_ids .shape )} "
703710 )
704711
712+ # print(f"{packed_text_ids.shape=}, {torch.sum(packed_text_ids, dtype=torch.float32)=}, {packed_text_ids=}")
705713 text_embeddings = self .language_model .embedding (
706714 input_ids = packed_text_ids , position_ids = packed_position_ids .unsqueeze (0 ),
707715 )
708- print (f"{ packed_text_ids .shape = } , { text_embeddings .shape = } " )
716+ # print(f"{packed_text_ids.shape=}, {text_embeddings.shape=}")
709717
710718 # Create the combined sequence tensor
711719 # For packed sequences: [total_seq, 1, hidden]
712720 combined_embeddings = text_embeddings .new_zeros (size = (sequence_length , self .language_model .config .hidden_size ))
713- print (f"{ packed_text_indexes .shape = } , { combined_embeddings .shape = } " )
721+ # print(f"{packed_text_indexes.shape=}, {combined_embeddings.shape=}")
714722 combined_embeddings [packed_text_indexes ] = text_embeddings .squeeze (1 )
723+ # print(f"{text_embeddings.shape=}, {torch.sum(text_embeddings, dtype=torch.float32)=}, {text_embeddings=}")
715724
716725 # --- Step 2: Process ViT tokens (understanding path) ---
717726 has_vision_path = (
@@ -733,18 +742,20 @@ def forward(
733742 cu_seqlens = cu_seqlens .to (torch .int32 )
734743 max_seqlen = torch .max (vit_token_seqlens ).item ()
735744 # [num_patches, 1, vit_hidden]
736- print (f"{ packed_vit_tokens .shape = } , { packed_vit_tokens .shape = } " )
745+ # print(f"{packed_vit_tokens.shape=}, {packed_vit_tokens.shape=}")
737746 image_embeddings = self ._embed_vision_tokens (
738747 pixel_values = packed_vit_tokens ,
739748 position_ids = packed_vit_position_ids ,
740749 cu_seqlens = cu_seqlens ,
741750 max_seqlen = max_seqlen ,
742751 )
743- print (f"{ image_embeddings .shape = } " )
752+ # print(f"after vit {image_embeddings.shape=}, {torch.sum(image_embeddings, dtype=torch.float32)=}, {image_embeddings=}")
753+ # print(f"{image_embeddings.shape=}")
744754
745755 if image_embeddings .dtype != combined_embeddings .dtype :
746756 image_embeddings = image_embeddings .to (combined_embeddings .dtype )
747757 combined_embeddings [packed_vit_token_indexes ] = image_embeddings .squeeze (1 )
758+ # print(f"{image_embeddings.shape=}, {torch.sum(image_embeddings, dtype=torch.float32)=}, {image_embeddings=}")
748759 elif has_vision_path :
749760 # Keep trainable vision modules in the autograd graph when this
750761 # microbatch contains no vision-understanding task.
@@ -771,6 +782,8 @@ def forward(
771782
772783 # --- Step 3: Process VAE latents (generation path) ---
773784 packed_latent_clean = None
785+ packed_latent_noise = None
786+ packed_timesteps_effective = None
774787 noise = None
775788 processed_timesteps = None
776789 has_generation_path = (
@@ -783,26 +796,52 @@ def forward(
783796 )
784797 has_gen_tokens = (
785798 has_generation_path
786- and padded_images is not None
799+ and (
800+ packed_latent_clean is not None
801+ or padded_latent is not None
802+ or padded_images is not None
803+ )
787804 and packed_vae_token_indexes is not None
788805 and packed_vae_token_indexes .numel () > 0
789806 and patchified_vae_latent_shapes is not None
790807 and packed_latent_position_ids is not None
791808 and packed_timesteps is not None
792809 )
793810 if has_gen_tokens :
794- vae_dtype = next (self .vae_model .parameters ()).dtype
795- padded_latent = self .vae_model .encode (padded_images .to (vae_dtype ))
796- packed_latent_clean = self ._patchify_vae_latents (
797- padded_latent ,
798- patchified_vae_latent_shapes ,
799- )
800- packed_latent , noise , processed_timesteps = (
801- self ._apply_flow_matching_noise (
802- packed_latent_clean ,
803- packed_timesteps ,
811+ if packed_latent_clean is None :
812+ if padded_latent is None :
813+ vae_dtype = next (self .vae_model .parameters ()).dtype
814+ padded_latent = self .vae_model .encode (padded_images .to (vae_dtype ))
815+ packed_latent_clean = self ._patchify_vae_latents (
816+ padded_latent ,
817+ patchified_vae_latent_shapes ,
818+ )
819+ if packed_latent_noise is not None :
820+ noise = packed_latent_noise .to (packed_latent_clean .dtype )
821+ if packed_timesteps_effective is None :
822+ processed_timesteps = torch .sigmoid (
823+ packed_timesteps .to (packed_latent_clean .dtype )
824+ )
825+ processed_timesteps = (
826+ self .timestep_shift
827+ * processed_timesteps
828+ / (1 + (self .timestep_shift - 1 ) * processed_timesteps )
829+ )
830+ else :
831+ processed_timesteps = packed_timesteps_effective .to (
832+ packed_latent_clean .dtype
833+ )
834+ packed_latent = (
835+ (1 - processed_timesteps [:, None ]) * packed_latent_clean
836+ + processed_timesteps [:, None ] * noise
837+ )
838+ else :
839+ packed_latent , noise , processed_timesteps = (
840+ self ._apply_flow_matching_noise (
841+ packed_latent_clean ,
842+ packed_timesteps ,
843+ )
804844 )
805- )
806845 packed_latent = self ._embed_generation_tokens (
807846 packed_latent ,
808847 processed_timesteps ,
@@ -811,6 +850,7 @@ def forward(
811850 combined_embeddings [packed_vae_token_indexes ] = packed_latent .to (
812851 combined_embeddings .dtype
813852 )
853+ # print(f"{packed_latent.shape=}, {torch.sum(packed_latent, dtype=torch.float32)=}, {packed_latent=}")
814854 elif has_generation_path :
815855 # Match the real generation path with one dummy image so the VAE
816856 # encoder and all generation-input modules stay in the graph.
@@ -875,7 +915,7 @@ def forward(
875915 attn_modes = attn_modes ,
876916 device = combined_embeddings .device ,
877917 )
878- print (f"{ packed_und_token_indexes .shape = } " )
918+ # print(f"{packed_und_token_indexes.shape=}")
879919
880920 # Run through GPTModel with decoder_input (skips embedding)
881921 hidden_states = self .language_model .decoder (
0 commit comments