Skip to content

Commit b4e5db2

Browse files
committed
add code for alignment
1 parent fcf9035 commit b4e5db2

3 files changed

Lines changed: 127 additions & 34 deletions

File tree

flagscale/models/megatron/bagel/models/bagel_model.py

Lines changed: 57 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -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(

flagscale/models/megatron/bagel/models/siglip_model.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -155,9 +155,9 @@ def forward(
155155
x = x + self.position_embeddings(packed_flattened_position_ids)
156156

157157
# Reshape for TransformerBlock: [seq_len, batch=1, hidden_size]
158-
x = x.unsqueeze(1)
158+
x = x.unsqueeze(1).contiguous()
159159

160-
print(f"{cu_seqlens=}, {max_seqlen=}")
160+
# print(f"{cu_seqlens=}, {max_seqlen=}")
161161
# Build PackedSeqParams for variable-length attention
162162
packed_seq_params = PackedSeqParams(
163163
cu_seqlens_q=cu_seqlens,
@@ -167,8 +167,9 @@ def forward(
167167
max_seqlen_kv=max_seqlen,
168168
)
169169

170-
print(f"{x.shape=}")
170+
# print(f"{x.shape=}")
171171
# Transformer forward
172+
# print(f"before encoder: {x.shape=}, {torch.sum(x, dtype=torch.float32)=}, {x=}, {packed_seq_params=}")
172173
x = self.decoder(
173174
hidden_states=x,
174175
attention_mask=None,

flagscale/train/megatron/train_bagel.py

Lines changed: 66 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
22
"""Pretrain BAGEL multimodal model for unified image understanding and generation."""
33

4+
import os
5+
46
import torch
57
from megatron.core import parallel_state
68
from megatron.core.enums import ModelType
@@ -129,6 +131,24 @@ def train_valid_test_datasets_provider(train_val_test_num_samples):
129131
return dataset_provider(train_val_test_num_samples)
130132

131133

134+
def _load_alignment_batch(path):
135+
"""Load and validate a CPU-only batch dumped by original BAGEL."""
136+
data = torch.load(path, map_location="cpu", weights_only=False)
137+
required = {
138+
"sequence_length", "sample_lens", "split_lens", "attn_modes",
139+
"packed_text_ids", "packed_text_indexes", "packed_position_ids",
140+
"packed_latent_clean", "packed_latent_noise",
141+
"packed_timesteps_effective",
142+
}
143+
missing = sorted(required.difference(data))
144+
if missing:
145+
raise ValueError(f"alignment batch is missing required fields: {missing}")
146+
metadata = data.get("_alignment_metadata", {})
147+
if metadata.get("source") != "original_bagel":
148+
raise ValueError(f"unexpected alignment batch source: {metadata!r}")
149+
return data
150+
151+
132152
def get_batch(data_iterator):
133153
"""Generate a batch from data iterator.
134154
@@ -151,8 +171,12 @@ def get_batch(data_iterator):
151171
assert (getattr(args, 'pipeline_model_parallel_size', 1) == 1)
152172
assert (getattr(args, 'tensor_model_parallel_size', 1) == 1)
153173

154-
data = next(data_iterator)
155-
print(f"{data=}")
174+
alignment_batch_path = os.getenv("BAGEL_ALIGN_BATCH")
175+
if alignment_batch_path:
176+
data = _load_alignment_batch(alignment_batch_path)
177+
else:
178+
data = next(data_iterator)
179+
# print(f"{data=}")
156180

157181
# --- Assemble batch dict ---
158182
batch = {
@@ -179,6 +203,14 @@ def get_batch(data_iterator):
179203
# --- VAE image generation fields ---
180204
if 'padded_images' in data:
181205
batch['padded_images'] = data['padded_images'].cuda(non_blocking=True)
206+
if 'padded_latent' in data:
207+
batch['padded_latent'] = data['padded_latent'].cuda(non_blocking=True)
208+
if 'packed_latent_clean' in data:
209+
batch['packed_latent_clean'] = data['packed_latent_clean'].cuda(non_blocking=True)
210+
if 'packed_latent_noise' in data:
211+
batch['packed_latent_noise'] = data['packed_latent_noise'].cuda(non_blocking=True)
212+
if 'packed_timesteps_effective' in data:
213+
batch['packed_timesteps_effective'] = data['packed_timesteps_effective'].cuda(non_blocking=True)
182214
if 'patchified_vae_latent_shapes' in data:
183215
batch['patchified_vae_latent_shapes'] = data['patchified_vae_latent_shapes']
184216
if 'packed_latent_position_ids' in data:
@@ -202,7 +234,7 @@ def get_batch(data_iterator):
202234
if 'mse_loss_indexes' in data:
203235
batch['mse_loss_indexes'] = data['mse_loss_indexes'].cuda(non_blocking=True)
204236

205-
print(f"{batch=}")
237+
# print(f"{batch=}")
206238
return batch
207239

208240

@@ -263,6 +295,10 @@ def forward_step(data_iterator, model: BagelModel):
263295
vit_token_seqlens=batch.get('vit_token_seqlens'),
264296
# VAE generation
265297
padded_images=batch.get('padded_images'),
298+
padded_latent=batch.get('padded_latent'),
299+
packed_latent_clean=batch.get('packed_latent_clean'),
300+
packed_latent_noise=batch.get('packed_latent_noise'),
301+
packed_timesteps_effective=batch.get('packed_timesteps_effective'),
266302
patchified_vae_latent_shapes=batch.get('patchified_vae_latent_shapes'),
267303
packed_latent_position_ids=batch.get('packed_latent_position_ids'),
268304
packed_vae_token_indexes=batch.get('packed_vae_token_indexes'),
@@ -319,20 +355,20 @@ def forward_step(data_iterator, model: BagelModel):
319355
mse = mse.mean(dim=-1).sum() * dp_world_size / total_mse_tokens
320356
combined_loss = combined_loss + mse * getattr(args, 'mse_weight', 1.0)
321357

322-
print(f"{ce=}, {mse=}")
358+
# print(f"{ce=}, {mse=}")
323359

324360
combined_loss = combined_loss + dummy
325361

326-
if torch.distributed.get_rank() == 0:
327-
print(
328-
"[BAGEL_LOSS_GRAPH]",
329-
f"ce_shape={None if ce is None else tuple(ce.shape)}",
330-
f"ce_grad_fn={None if ce is None else ce.grad_fn}",
331-
f"mse_shape={None if mse is None else tuple(mse.shape)}",
332-
f"mse_grad_fn={None if mse is None else mse.grad_fn}",
333-
f"combined_grad_fn={combined_loss.grad_fn}",
334-
flush=True,
335-
)
362+
# if torch.distributed.get_rank() == 0:
363+
# print(
364+
# "[BAGEL_LOSS_GRAPH]",
365+
# f"ce_shape={None if ce is None else tuple(ce.shape)}",
366+
# f"ce_grad_fn={None if ce is None else ce.grad_fn}",
367+
# f"mse_shape={None if mse is None else tuple(mse.shape)}",
368+
# f"mse_grad_fn={None if mse is None else mse.grad_fn}",
369+
# f"combined_grad_fn={combined_loss.grad_fn}",
370+
# flush=True,
371+
# )
336372

337373
return combined_loss.unsqueeze(0), bagel_loss_func
338374

@@ -411,6 +447,22 @@ def add_bagel_extra_args(parser):
411447
help='Reweight CE loss using per-token ce_loss_weights.',
412448
)
413449

450+
# Random preprocessing controls. Keep these explicit so alignment runs do
451+
# not silently fall back to BagelDataConfig defaults.
452+
group.add_argument(
453+
'--text-cond-dropout-prob', type=float, default=0.0,
454+
help='Probability of dropping text conditioning during packing.',
455+
)
456+
group.add_argument(
457+
'--vit-cond-dropout-prob', type=float, default=0.0,
458+
help='Probability of dropping ViT conditioning during packing.',
459+
)
460+
group.add_argument(
461+
'--vae-cond-dropout-prob', type=float, default=0.0,
462+
help='Probability of dropping VAE conditioning during packing.',
463+
)
464+
465+
414466
# Freeze controls
415467
group.add_argument(
416468
'--freeze-VAE', action='store_true', default=False,

0 commit comments

Comments
 (0)