Skip to content

[ckpt] fix: align packed tensor offsets to dtype itemsize for mixed-dtype streams - #7318

Open
savaresejeremy wants to merge 1 commit into
verl-project:mainfrom
savaresejeremy:fix/ckpt-packing-alignment
Open

savaresejeremy wants to merge 1 commit into
verl-project:mainfrom
savaresejeremy:fix/ckpt-packing-alignment

Conversation

@savaresejeremy

@savaresejeremy savaresejeremy commented Aug 7, 2026 •

Copy link
Copy Markdown

What does this PR do?

Fixes a latent crash in the checkpoint-engine wire for mixed-dtype weight streams.

The nccl/hccl/nixl checkpoint engines pack tensors back-to-back into a uint8 bucket and record each start in TensorMeta.offset. On receive, each single-chunk tensor is materialized with Tensor.view(dtype) on a slice of the bucket, and torch requires that slice's storage offset to be a multiple of the dtype's element size. A homogeneous stream keeps that invariant for free (every tensor's byte length is a multiple of its own element size), which is why today's bf16-only exports never hit it. A mixed-dtype stream breaks it: an odd-numel bf16 tensor (6 bytes) followed by an fp32 tensor puts the fp32 start at offset % 4 == 2, and the receiver raises at the view. The fix lives in the engines rather than a backend because the wire should not crash on any legal stream, whichever of FSDP, Megatron or veomni produces it.

The fix is sender-side only: a shared align_bucket_offset helper rounds each tensor's bucket offset up to its dtype's itemsize at pack time. The receiver already slices strictly by the recorded offsets, so no receive-path change is needed, and homogeneous streams produce a byte-identical wire layout to before (the alignment is a no-op there). On mixed streams the cost is bounded padding, at most itemsize minus one bytes ahead of a tensor whose start would otherwise misalign; pad gaps are never read.

This is worth fixing now rather than when it fires:

  • The open [fsdp] fix: preserve Hugging Face fp32-keep modules in FSDP2 #7165 restores Hugging Face's fp32-keep contract (under the default bf16 policy, the _keep_in_fp32_modules_strict declarations); once it lands, models with strict fp32-keep modules export mixed bf16+fp32 streams down the default FSDP2 path.
  • Quantized weight refit produces mixed streams by construction: the in-tree QAT export (verl/utils/qat/quantizer.py, wired into get_per_tensor_param behind the QAT config) already yields packed narrow weights alongside their scales and bf16 passthrough tensors, and the very large models discussed in discussion 7273 push the same (weight, scale) shape through whichever engine carries them.

#7263, our open checkpoint-engine PR, notes this residual case in its text; its 16-byte rounding fixes the bucket-cut variant in the engine it adds, and this PR fixes the per-tensor packing case shared by the existing engines.

Checklist Before Starting

Reproduction

On a machine with two GPUs at main, drive the stock engine with a stream that puts an odd-numel bf16 tensor immediately ahead of an fp32 tensor; the crash needs the running byte offset in front of a tensor to not be a multiple of its element size, so what matters is the total byte length of everything packed before it. The receiver crashes at the view:

# repro_mixed_dtype.py: drive NCCLCheckpointEngine with a mixed-dtype inventory.
# Runs on one node with 2 GPUs. Exit 0 = transfer completed and every received
# tensor compared bit-equal; nonzero = the receiver raised (prints the error).
import asyncio
import sys

import ray
import torch


@ray.remote(num_gpus=1)
class Sender:
    def __init__(self, bucket_mb: int):
        from verl.checkpoint_engine.nccl_checkpoint_engine import NCCLCheckpointEngine

        self.engine = NCCLCheckpointEngine(bucket_size=bucket_mb << 20, is_master=True)
        self.meta = self.engine.prepare()

    def metadata(self):
        return {"zmq_ip": self.meta.zmq_ip, "zmq_port": self.meta.zmq_port}

    def init_group(self, world_size: int):
        self.engine.init_process_group(rank=0, world_size=world_size, master_metadata=self.meta)

    def send(self):
        def inventory():
            g = torch.Generator().manual_seed(7)
            # odd-numel bf16 (6 bytes) ahead of fp32: the misalignment trigger
            yield "norm.weight", torch.randn(3, generator=g).to(torch.bfloat16).cuda()
            yield "router.weight", torch.randn(64, 64, generator=g).cuda()  # fp32
            yield "proj.weight", torch.randn(4096, 512, generator=g).to(torch.bfloat16).cuda()
            yield "step", torch.arange(3, dtype=torch.int64).cuda()

        asyncio.run(self.engine.send_weights(inventory()))
        return True


@ray.remote(num_gpus=1)
class Receiver:
    def __init__(self, bucket_mb: int):
        from verl.checkpoint_engine.nccl_checkpoint_engine import NCCLCheckpointEngine

        self.engine = NCCLCheckpointEngine(bucket_size=bucket_mb << 20, is_master=False)
        self.engine.prepare()

    def init_group(self, world_size: int, master_metadata: dict):
        from verl.checkpoint_engine.nccl_checkpoint_engine import MasterMetadata

        self.engine.init_process_group(
            rank=1, world_size=world_size, master_metadata=MasterMetadata(**master_metadata)
        )

    def receive_and_check(self):
        async def run():
            received = {}
            async for name, tensor in self.engine.receive_weights():
                received[name] = tensor.clone()
            return received

        received = asyncio.run(run())
        g = torch.Generator().manual_seed(7)
        expected = {
            "norm.weight": torch.randn(3, generator=g).to(torch.bfloat16),
            "router.weight": torch.randn(64, 64, generator=g),
            "proj.weight": torch.randn(4096, 512, generator=g).to(torch.bfloat16),
            "step": torch.arange(3, dtype=torch.int64),
        }
        for name, exp in expected.items():
            got = received[name].cpu()
            assert got.dtype == exp.dtype and got.shape == exp.shape, name
            assert torch.equal(got.view(torch.uint8).view(-1), exp.view(torch.uint8).view(-1)), name
        return sorted(received)


def main():
    ray.init(num_gpus=2)
    try:
        sender = Sender.remote(bucket_mb=128)
        receiver = Receiver.remote(bucket_mb=128)
        meta = ray.get(sender.metadata.remote())
        ray.get([sender.init_group.remote(2), receiver.init_group.remote(2, meta)])
        send_ref = sender.send.remote()
        names = ray.get(receiver.receive_and_check.remote())
        ray.get(send_ref)
        print(f"OK: transfer completed, {len(names)} tensors bit-equal: {names}")
        return 0
    finally:
        ray.shutdown()


if __name__ == "__main__":
    sys.exit(main())
ray.exceptions.RayTaskError(RuntimeError): ray::Receiver.receive_and_check()
  File ".../repro_mixed_dtype.py", line 57, in run
    async for name, tensor in self.engine.receive_weights():
  File ".../verl/checkpoint_engine/nccl_checkpoint_engine.py", line 318, in receive_weights
    async for name, weight in merge_weight_chunks(self._receive_weight_chunks(), self.bucket_size):
  File ".../verl/checkpoint_engine/base.py", line 604, in merge_weight_chunks
    name, weight = tensor_meta.name, chunk.view(tensor_meta.dtype).view(tensor_meta.shape)
RuntimeError: self.storage_offset() must be divisible by 4 to view Byte as Float (different element sizes), but got 6

With this PR applied, the same run completes: OK: transfer completed, 4 tensors bit-equal: ['norm.weight', 'proj.weight', 'router.weight', 'step'].

Test

CPU lane (no GPU; runs in cpu_unit_tests):

pytest tests/checkpoint_engine/test_packing_alignment_on_cpu.py -q
  • align_bucket_offset math across dtypes (uint8 through complex128, no-op when aligned, no-op at bucket start).
  • A mixed-dtype inventory driven through the real split_weight_chunks and receive-side merge_weight_chunks, with a pack simulation mirroring the senders' loop, round-trips bit-exactly across bucket sizes that force flushes, exact bucket fills, multi-chunk continuation, and flushes triggered by the alignment padding itself (70/101/256/257/300/512 bytes and 1 MiB), with every tensor-start offset checked against its itemsize.
  • Control: the same inventory packed back-to-back (the pre-fix layout) is rejected at the view call with the divisibility error, proving the round-trip test can fail.
  • A homogeneous bf16 stream produces identical bucket counts, offsets and bucket bytes with and without the fix, at flush-forcing and single-bucket sizes (wire layout unchanged for all current users).
  • A source-level guard asserts every engine's packing loop calls the helper, since the engines' send loops need live transports and cannot run in this lane.

Results: 19 passed in 14.67s (CPU, inside the training image used for the runs below).

GPU regression check (stock engine unaffected on the homogeneous path):

pytest tests/checkpoint_engine/test_correctness_on_gpu.py -q

Results: test_nccl_checkpoint_engine ran in both rebuild_group parametrizations on one 8-GPU B200 node (2 trainer + 6 rollout workers, 3 update rounds each, weight comparison on) with the model path overridden to a locally staged Qwen2.5-0.5B-Instruct; both passed. The file's nixl and NPU cases are upstream-skipped and did not run.

End-to-end A/B (same recipe, same injected mixed-dtype stream, only this fix differs):

arm tree result
A (stock) main @ 474c2f4 + inject shim crashed at the first weight sync with RuntimeError: self.storage_offset() must be divisible by 4 to view Byte as Float (different element sizes), but got 6 at merge_weight_chunks (base.py:604); the job's single automatic restart also failed and the job finished FAILED
B (fixed) main @ 474c2f4 + inject shim + this PR completed all 5 steps; weight syncs 3.3-4.3 s; job finished SUCCEEDED

The injection shim is an env-gated test hook (identical commit in both arms) that prepends one odd-numel bf16 tensor and one fp32 tensor to the export stream, emulating the stream shape #7165 and quantized refit will produce. Each arm ran the same GRPO one-step-off recipe on two 8-GPU B200 nodes (one trainer node, one rollout node), Qwen3.5-27B bf16, rollout TP=4, 2048 MB buckets, 5 training steps, identical seeds.

Validation scope: nccl engine end-to-end + backend-independent CPU tests of the shared pack/merge path; the hccl and nixl senders carry the same alignment call (hccl is exercised by the NPU CI lane). The mooncake engine casts to the rollout dtype before packing and the kimi engine takes its offsets from the external checkpoint-engine package, so neither packs mixed-dtype buckets this way.

Additional Info

  • The colocated vLLM update path (verl/workers/rollout/vllm_rollout/bucketed_weight_transfer.py) has the same back-to-back packing and we are happy to bring the same fix there as a follow-up.
  • Sequencing with [ckpt] feat: add nccl_parallel checkpoint engine (all actor ranks send) #7263: that PR (in review) adds a fourth engine with its own packing loop and publicly flagged this issue in its text; once both are in-tree it adopts the same alignment call in its loop as a follow-through.
  • AI disclosure: this fix was developed with AI assistance (Claude, via internal agent tooling). All changes were human-reviewed and the branch was reviewed internally before submission; the commit carries a Co-authored-by: Claude trailer per AGENTS.md.

API and Usage Example

No API, CLI or config change; behavior is identical for homogeneous-dtype streams.

Design & Code Changes

  • verl/checkpoint_engine/base.py: new align_bucket_offset(offset, dtype) helper next to the packing utilities.
  • verl/checkpoint_engine/nccl_checkpoint_engine.py, hccl_checkpoint_engine.py, nixl_checkpoint_engine.py: one alignment call at the top of each send loop, before the bucket-overflow check.
  • tests/checkpoint_engine/test_packing_alignment_on_cpu.py: new CPU-lane tests described above.

…type streams

Bucket packing in the nccl/hccl/nixl checkpoint engines places tensors
back-to-back in a uint8 buffer. The receive path reinterprets each
single-chunk slice with Tensor.view(dtype), which torch only permits when
the slice's storage offset is a multiple of the element size, so a
mixed-dtype stream (for example an odd-numel bf16 tensor followed by an
fp32 tensor) crashes the receiver. Homogeneous streams are unaffected
because every tensor's byte length is a multiple of its own element size.

Fix: a shared align_bucket_offset helper in base.py rounds each tensor's
bucket offset up to its dtype's itemsize at pack time. The receiver
already slices strictly by the offsets recorded in TensorMeta, so the
change is sender-side only and the wire layout of homogeneous streams is
byte-identical to before.

Tests (CPU lane): helper math across dtypes; a mixed-dtype inventory
driven through the real split_weight_chunks and merge_weight_chunks with a
pack simulation mirroring the senders' loop, across bucket sizes that
force flushes, exact fills, multi-chunk continuation, and flushes
triggered by the alignment padding itself; a control showing the pre-fix
back-to-back layout is rejected at the view call; homogeneous-stream
bucket-structure no-op checks at flush-forcing and single-bucket sizes;
and a source-level guard, its file list derived from the packing-loop
signature, asserting every such loop calls the helper.

Co-authored-by: Claude
@wuxibin89

Copy link
Copy Markdown
Collaborator

The condition of latent crash is pretty strict, I ran DeepSeek V4 Flash with fp8/fp4 weight and scale exported, didn't trigger crash:

examples/grpo_trainer/run_deepseek_v4_veomni.sh

@savaresejeremy

Copy link
Copy Markdown
Author

You are right, and honestly this one is unlikely to pop up on its own. We included it because we hit a similar crash during #7263's 4-node validation and thought it prudent to cover this one for completeness. That one occurred naturally because #7263 divides the bucket budget across senders, and at 24 senders the per-sender bucket size comes out odd, so a bucket cut through a large tensor left everything behind it misaligned (fixed there by rounding bucket sizes). Unlike that one, this one is unlikely to occur with modern architectures: their dimensions are even, so every tensor's byte length is a multiple of 4 and back-to-back packing stays aligned on its own. It would take an odd byte length ahead of a wider dtype to break, and nothing in-tree has that shape today.

Your run additionally could not hit it for two reasons: fp8/fp4 payloads have element size 1, so they never misalign themselves, and V4's even dimensions keep the running offset a multiple of 4 in front of every fp32 scale (also, run_deepseek_v4_veomni.sh syncs through the colocated bucketed_weight_transfer.py rather than the checkpoint-engine wire this PR touches. That path packs the same way and is the follow-up named in Additional Info).

If you do want to see it fire on a real run, apply the small env-gated patch below: it prepends one odd-numel bf16 tensor and one fp32 tensor to the export stream and run the stock one-step-off script with the nccl engine.

git apply mixed_dtype_inject.patch
MODEL_PATH=Qwen/Qwen3-0.6B VERL_AB_INJECT_MIXED_DTYPE=1 \
bash verl/experimental/one_step_off_policy/shell/grpo_0.6b_gsm8k_fsdp2_2_6.sh \
  actor_rollout_ref.rollout.checkpoint_engine.backend=nccl

On main this crashes at the first weight sync with RuntimeError: self.storage_offset() must be divisible by 4 to view Byte as Float (different element sizes), but got 6; with this PR it trains through.

mixed_dtype_inject.patch (~36 lines, test hook only)
diff --git a/verl/checkpoint_engine/base.py b/verl/checkpoint_engine/base.py
index 3c0d230..7e1e15c 100644
--- a/verl/checkpoint_engine/base.py
+++ b/verl/checkpoint_engine/base.py
@@ -334,6 +334,20 @@ class CheckpointEngineWorker(Worker):
     @register(dispatch_mode=Dispatch.ONE_TO_ALL, blocking=False)
     async def update_weights(self, global_steps: int = None):
         weights = self.checkpoint_engine.receive_weights(global_steps=global_steps)
+        # --- A/B TEST SHIM (never merge): drop the injected fake tensors after
+        # the wire (past the crash site under test) so the rollout loader never
+        # sees unknown parameter names. Same env gate as the sender shim.
+        import os as _os
+
+        if _os.getenv("VERL_AB_INJECT_MIXED_DTYPE") == "1":
+            _wire = weights
+
+            async def _drop_injected():
+                async for name, tensor in _wire:
+                    if not str(name).startswith("_ab_inject."):
+                        yield name, tensor
+
+            weights = _drop_injected()
         await self.server_adapter.update_weights(
             weights,
             global_steps=global_steps,
diff --git a/verl/workers/engine/fsdp/transformer_impl.py b/verl/workers/engine/fsdp/transformer_impl.py
index dd5cf33..8cfd86f 100644
--- a/verl/workers/engine/fsdp/transformer_impl.py
+++ b/verl/workers/engine/fsdp/transformer_impl.py
@@ -947,6 +947,28 @@ class FSDPEngine(BaseEngine):
         return hf_delta_export(gen, self._delta_shard_snap, self._hf_delta_entry), None
 
     def get_per_tensor_param(self, layered_summon=False, base_sync_done=False, **kwargs):
+        # --- A/B TEST SHIM (never merge): emulate a mixed-dtype export stream.
+        # Gated on VERL_AB_INJECT_MIXED_DTYPE=1; prepends one odd-numel bf16
+        # tensor (6 bytes) then one fp32 tensor, the stream shape #7165
+        # (fp32-keep modules) and quantized (weight, scale) refit produce.
+        # This commit is byte-identical in both A/B arms; the arms differ
+        # only by the alignment-fix commit under test.
+        import os as _os
+
+        if _os.getenv("VERL_AB_INJECT_MIXED_DTYPE") == "1" and not kwargs.pop("_ab_no_inject", False):
+            base_gen, base_peft = self.get_per_tensor_param(
+                layered_summon=layered_summon, base_sync_done=base_sync_done, _ab_no_inject=True, **kwargs
+            )
+
+            def _injected():
+                import torch as _torch
+
+                yield "_ab_inject.norm_bf16_odd", _torch.full((3,), 1.0, dtype=_torch.bfloat16, device="cuda")
+                yield "_ab_inject.scale_fp32", _torch.full((4, 4), 2.0, dtype=_torch.float32, device="cuda")
+                yield from base_gen
+
+            return _injected(), base_peft
+        kwargs.pop("_ab_no_inject", None)
         log_gpu_memory_usage("Before load_fsdp_model_to_gpu", logger=logger)
 
         # FSDP2 CPUOffloadPolicy owns CPU<->GPU placement; calling model.to(device) here

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants