CoolFace
Modelpublic

Motif-Technologies/optimizer

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
58likes235downloads
test_cpu_offload.py967 linesDownload Raw Back to test
1"""CPU offloading tests for optimizer states.2 3Run with:4    torchrun --nproc-per-node=8 --local-ranks-filter=0 test/test_cpu_offload.py5 6Tests:7    1. Correctness: turn_on_cpu_offload() produces identical results to no offload8    2. Memory: GPU optimizer state storage is freed after offload9    3. AdamW: moment1/moment2 offloading works correctly10"""11 12import copy13import logging14 15import pytest16import torch17import torch.distributed as dist18from torch.distributed.tensor import DTensor, Shard, distribute_tensor19 20logger = logging.getLogger(__name__)21logging.basicConfig(level=logging.INFO, format="[%(levelname)s] %(message)s")22 23 24def _setup():25    dist.init_process_group(backend="nccl")26    rank = dist.get_rank()27    torch.cuda.set_device(rank % torch.cuda.device_count())28    return rank, dist.get_world_size()29 30 31def _make_mesh(world_size):32    return dist.init_device_mesh("cuda", (world_size, ),33                                 mesh_dim_names=("dp", ))34 35 36def test_correctness(rank, world_size):37    """Verify that turn_on_cpu_offload() produces identical parameters as no offload."""38    from optimizer.muon import Muon39    from optimizer.newton_schulz import set_ns_compile40 41    set_ns_compile(False)42    torch.manual_seed(42)43 44    mesh = _make_mesh(world_size)45 46    dim0, dim1 = 64, 12847    num_params = 448    num_steps = 349 50    # Pre-generate all data on all ranks (same seed โ†’ same values).51    full_params = [52        torch.randn(dim0, dim1, device="cuda") for _ in range(num_params)53    ]54    full_grads = [[55        torch.randn(dim0, dim1, device="cuda") for _ in range(num_params)56    ] for _ in range(num_steps)]57 58    def make_optimizer(cpu_offload):59        params, names = [], []60        for i, fp in enumerate(full_params):61            dt = distribute_tensor(fp.clone(), mesh, [Shard(0)])62            p = torch.nn.Parameter(dt)63            params.append(p)64            names.append(f"layer.{i}.weight")65        param_groups = [{66            "params": params,67            "names": names,68            "use_muon": True,69            "lr": 0.02,70            "weight_decay": 0.01,71            "momentum": 0.95,72            "nesterov": True,73            "ns_steps": 5,74            "none_grad": False,75        }]76        optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)77        if cpu_offload:78            optim.turn_on_cpu_offload()79        return optim, params80 81    optim_ref, params_ref = make_optimizer(False)82    optim_off, params_off = make_optimizer(True)83 84    for step_idx in range(num_steps):85        for i in range(num_params):86            g = full_grads[step_idx][i]87            params_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])88            params_off[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])89 90        optim_ref.step()91        optim_off.step()92 93        for i in range(num_params):94            ref_full = params_ref[i].data.full_tensor()95            off_full = params_off[i].data.full_tensor()96            torch.testing.assert_close(ref_full, off_full, atol=0, rtol=0)97 98        if rank == 0:99            logger.info("Step %d: correctness OK", step_idx)100 101    set_ns_compile(True)102    if rank == 0:103        logger.info("PASSED: test_correctness")104 105 106def test_memory(rank, world_size):107    """Verify that GPU storage is freed after offload."""108    from optimizer.muon import Muon109    from optimizer.newton_schulz import set_ns_compile110 111    set_ns_compile(False)112    torch.manual_seed(42)113 114    mesh = _make_mesh(world_size)115 116    dim0, dim1 = 512, 1024117    num_params = 8118 119    params, names = [], []120    for i in range(num_params):121        full = torch.randn(dim0, dim1, device="cuda")122        dt = distribute_tensor(full, mesh, [Shard(0)])123        p = torch.nn.Parameter(dt)124        p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),125                                   mesh, [Shard(0)])126        params.append(p)127        names.append(f"layer.{i}.weight")128 129    param_groups = [{130        "params": params,131        "names": names,132        "use_muon": True,133        "lr": 0.02,134        "weight_decay": 0.01,135        "momentum": 0.95,136        "nesterov": True,137        "ns_steps": 5,138        "none_grad": False,139    }]140    optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)141    optim.turn_on_cpu_offload()142 143    optim.step()144    torch.cuda.synchronize()145 146    # After step + offload, all momentum buffer GPU storage should be freed.147    for p in params:148        state = optim.state[p]149        if "momentum_buffer" not in state:150            continue151        buf = state["momentum_buffer"]152        local_buf = buf._local_tensor if isinstance(buf, DTensor) else buf153        assert local_buf.untyped_storage().size() == 0, (154            f"Expected freed GPU storage after offload, got "155            f"{local_buf.untyped_storage().size()} bytes")156 157    # Verify CPU pool has pinned buffers.158    pool = optim._cpu_offload_pool159    assert len(pool._managed) > 0, "No tensors tracked by CPU offload pool"160    for grp in pool._groups.values():161        assert grp["cpu_flat"].is_pinned(), "CPU buffer must be pinned memory"162 163    # Run another step to verify reload + compute + offload cycle works.164    for p in params:165        p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),166                                   mesh, [Shard(0)])167    optim.step()168    torch.cuda.synchronize()169 170    # Storage should be freed again after second step.171    for p in params:172        state = optim.state[p]173        if "momentum_buffer" not in state:174            continue175        buf = state["momentum_buffer"]176        local_buf = buf._local_tensor if isinstance(buf, DTensor) else buf177        assert local_buf.untyped_storage().size() == 0178 179    set_ns_compile(True)180    if rank == 0:181        logger.info("PASSED: test_memory")182 183 184def test_adamw_offload(rank, world_size):185    """Verify AdamW moment1/moment2 are offloaded correctly."""186    from optimizer.muon import Muon187    from optimizer.newton_schulz import set_ns_compile188 189    set_ns_compile(False)190    torch.manual_seed(42)191 192    mesh = _make_mesh(world_size)193 194    num_steps = 3195 196    # Create both Muon (2D) and AdamW (1D) params.197    muon_params, muon_names = [], []198    adamw_params, adamw_names = [], []199 200    for i in range(4):201        full = torch.randn(64, 128, device="cuda")202        dt = distribute_tensor(full, mesh, [Shard(0)])203        p = torch.nn.Parameter(dt)204        muon_params.append(p)205        muon_names.append(f"layer.{i}.weight")206 207    for i in range(3):208        full = torch.randn(128, device="cuda")209        dt = distribute_tensor(full, mesh, [Shard(0)])210        p = torch.nn.Parameter(dt)211        adamw_params.append(p)212        adamw_names.append(f"layer.{i}.bias")213 214    # Pre-generate grads.215    muon_grads = [[torch.randn(64, 128, device="cuda") for _ in range(4)]216                  for _ in range(num_steps)]217    adamw_grads = [[torch.randn(128, device="cuda") for _ in range(3)]218                   for _ in range(num_steps)]219 220    def make_optimizer(cpu_offload):221        mp = [222            torch.nn.Parameter(223                distribute_tensor(p.data.full_tensor().clone(), mesh,224                                  [Shard(0)])) for p in muon_params225        ]226        ap = [227            torch.nn.Parameter(228                distribute_tensor(p.data.full_tensor().clone(), mesh,229                                  [Shard(0)])) for p in adamw_params230        ]231        param_groups = [232            {233                "params": mp,234                "names": list(muon_names),235                "use_muon": True,236                "lr": 0.02,237                "weight_decay": 0.01,238                "momentum": 0.95,239                "nesterov": True,240                "ns_steps": 5,241                "none_grad": False,242                "adamw_betas": (0.9, 0.95),243                "adamw_eps": 1e-8,244            },245            {246                "params": ap,247                "use_muon": False,248                "lr": 1e-3,249                "weight_decay": 0.01,250                "adamw_betas": (0.9, 0.95),251                "adamw_eps": 1e-8,252            },253        ]254        optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)255        if cpu_offload:256            optim.turn_on_cpu_offload()257        return optim, mp, ap258 259    optim_ref, mp_ref, ap_ref = make_optimizer(False)260    optim_off, mp_off, ap_off = make_optimizer(True)261 262    for step_idx in range(num_steps):263        for i in range(4):264            g = muon_grads[step_idx][i]265            mp_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])266            mp_off[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])267        for i in range(3):268            g = adamw_grads[step_idx][i]269            ap_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])270            ap_off[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])271 272        optim_ref.step()273        optim_off.step()274 275        # Compare Muon params.276        for i in range(4):277            ref_full = mp_ref[i].data.full_tensor()278            off_full = mp_off[i].data.full_tensor()279            torch.testing.assert_close(ref_full, off_full, atol=0, rtol=0)280 281        # Compare AdamW params.282        for i in range(3):283            ref_full = ap_ref[i].data.full_tensor()284            off_full = ap_off[i].data.full_tensor()285            torch.testing.assert_close(ref_full, off_full, atol=0, rtol=0)286 287        if rank == 0:288            logger.info("Step %d: AdamW offload correctness OK", step_idx)289 290    # Verify AdamW states are offloaded.291    for p in ap_off:292        state = optim_off.state[p]293        for key in ("moment1", "moment2"):294            if key not in state:295                continue296            t = state[key]297            local_t = t._local_tensor if isinstance(t, DTensor) else t298            assert local_t.untyped_storage().size() == 0, (299                f"AdamW {key} storage not freed after offload")300 301    set_ns_compile(True)302    if rank == 0:303        logger.info("PASSED: test_adamw_offload")304 305 306def test_memory_savings(rank, world_size):307    """Measure actual GPU memory difference with and without offload."""308    from optimizer.muon import Muon309    from optimizer.newton_schulz import set_ns_compile310 311    set_ns_compile(False)312 313    mesh = _make_mesh(world_size)314    dim0, dim1 = 1024, 2048315    num_params = 8316 317    def run_step(cpu_offload):318        torch.cuda.empty_cache()319        torch.cuda.reset_peak_memory_stats()320        torch.manual_seed(42)321 322        params, names = [], []323        for i in range(num_params):324            full = torch.randn(dim0, dim1, device="cuda")325            dt = distribute_tensor(full, mesh, [Shard(0)])326            p = torch.nn.Parameter(dt)327            p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),328                                       mesh, [Shard(0)])329            params.append(p)330            names.append(f"layer.{i}.weight")331 332        param_groups = [{333            "params": params,334            "names": names,335            "use_muon": True,336            "lr": 0.02,337            "weight_decay": 0.01,338            "momentum": 0.95,339            "nesterov": True,340            "ns_steps": 5,341            "none_grad": False,342        }]343        optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)344        if cpu_offload:345            optim.turn_on_cpu_offload()346        optim.step()347        torch.cuda.synchronize()348 349        mem = torch.cuda.memory_allocated()350        # Clean up to avoid interference.351        del optim, params, param_groups352        torch.cuda.empty_cache()353        return mem354 355    mem_no_offload = run_step(False)356    mem_with_offload = run_step(True)357 358    if rank == 0:359        logger.info("Memory without offload: %.2f MB",360                    mem_no_offload / 1024**2)361        logger.info("Memory with offload:    %.2f MB",362                    mem_with_offload / 1024**2)363        saved = mem_no_offload - mem_with_offload364        logger.info("Memory saved:           %.2f MB", saved / 1024**2)365 366    assert mem_with_offload < mem_no_offload, (367        f"Expected memory reduction with CPU offload. "368        f"Without: {mem_no_offload / 1024**2:.2f} MB, "369        f"With: {mem_with_offload / 1024**2:.2f} MB")370 371    set_ns_compile(True)372    if rank == 0:373        logger.info("PASSED: test_memory_savings")374 375 376def test_toggle_correctness(rank, world_size):377    """Verify toggling offload on/off between steps produces identical results."""378    from optimizer.muon import Muon379    from optimizer.newton_schulz import set_ns_compile380 381    set_ns_compile(False)382    torch.manual_seed(42)383 384    mesh = _make_mesh(world_size)385 386    dim0, dim1 = 64, 128387    num_params = 4388    num_steps = 6389 390    full_params = [391        torch.randn(dim0, dim1, device="cuda") for _ in range(num_params)392    ]393    full_grads = [[394        torch.randn(dim0, dim1, device="cuda") for _ in range(num_params)395    ] for _ in range(num_steps)]396 397    def make_optimizer():398        params, names = [], []399        for i, fp in enumerate(full_params):400            dt = distribute_tensor(fp.clone(), mesh, [Shard(0)])401            p = torch.nn.Parameter(dt)402            params.append(p)403            names.append(f"layer.{i}.weight")404        param_groups = [{405            "params": params,406            "names": names,407            "use_muon": True,408            "lr": 0.02,409            "weight_decay": 0.01,410            "momentum": 0.95,411            "nesterov": True,412            "ns_steps": 5,413            "none_grad": False,414        }]415        optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)416        return optim, params417 418    # Reference: no offload at all.419    optim_ref, params_ref = make_optimizer()420 421    # Toggle: on โ†’ step โ†’ off โ†’ step โ†’ on โ†’ step ...422    optim_toggle, params_toggle = make_optimizer()423 424    for step_idx in range(num_steps):425        # Toggle offload every 2 steps: on for [0,1], off for [2,3], on for [4,5].426        want_on = (step_idx // 2) % 2 == 0427        if want_on and not optim_toggle.cpu_offload:428            optim_toggle.turn_on_cpu_offload()429        elif not want_on and optim_toggle.cpu_offload:430            optim_toggle.turn_off_cpu_offload()431 432        for i in range(num_params):433            g = full_grads[step_idx][i]434            params_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])435            params_toggle[i].grad = distribute_tensor(g.clone(), mesh,436                                                      [Shard(0)])437 438        optim_ref.step()439        optim_toggle.step()440 441        for i in range(num_params):442            ref_full = params_ref[i].data.full_tensor()443            tog_full = params_toggle[i].data.full_tensor()444            torch.testing.assert_close(ref_full, tog_full, atol=0, rtol=0)445 446        if rank == 0:447            logger.info(448                "Step %d (offload=%s): toggle correctness OK",449                step_idx,450                optim_toggle.cpu_offload,451            )452 453    set_ns_compile(True)454    if rank == 0:455        logger.info("PASSED: test_toggle_correctness")456 457 458def test_leak(rank, world_size):459    """Run many iterations and verify no CPU/GPU memory leak."""460    import os461 462    from optimizer.muon import Muon463    from optimizer.newton_schulz import set_ns_compile464 465    set_ns_compile(False)466    torch.manual_seed(42)467 468    mesh = _make_mesh(world_size)469 470    dim0, dim1 = 512, 1024471    num_params = 8472    num_steps = 50473 474    params, names = [], []475    for i in range(num_params):476        full = torch.randn(dim0, dim1, device="cuda")477        dt = distribute_tensor(full, mesh, [Shard(0)])478        p = torch.nn.Parameter(dt)479        params.append(p)480        names.append(f"layer.{i}.weight")481 482    param_groups = [{483        "params": params,484        "names": names,485        "use_muon": True,486        "lr": 0.02,487        "weight_decay": 0.01,488        "momentum": 0.95,489        "nesterov": True,490        "ns_steps": 5,491        "none_grad": False,492    }]493    optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)494    optim.turn_on_cpu_offload()495 496    def get_cpu_rss_mb():497        """Get current process RSS in MB from /proc/self/statm."""498        with open("/proc/self/statm") as f:499            pages = int(f.read().split()[1])500        return pages * os.sysconf("SC_PAGE_SIZE") / (1024**2)501 502    gpu_after_warmup = None503    cpu_after_warmup = None504 505    for step_idx in range(num_steps):506        for p in params:507            p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),508                                       mesh, [Shard(0)])509 510        optim.step()511        torch.cuda.synchronize()512 513        gpu_mem = torch.cuda.memory_allocated()514        cpu_mem = get_cpu_rss_mb()515 516        # Record baseline after warmup (step 2 โ€” first step creates states,517        # second step does first full offload/reload cycle).518        if step_idx == 2:519            gpu_after_warmup = gpu_mem520            cpu_after_warmup = cpu_mem521 522        if rank == 0 and step_idx % 10 == 0:523            logger.info(524                "Step %d: GPU alloc=%.2f MB, CPU RSS=%.2f MB",525                step_idx,526                gpu_mem / (1024**2),527                cpu_mem,528            )529 530    # Final measurements.531    torch.cuda.synchronize()532    gpu_final = torch.cuda.memory_allocated()533    cpu_final = get_cpu_rss_mb()534 535    if rank == 0:536        logger.info(537            "After %d steps: GPU alloc=%.2f MB, CPU RSS=%.2f MB",538            num_steps,539            gpu_final / (1024**2),540            cpu_final,541        )542        logger.info(543            "Warmup baseline: GPU alloc=%.2f MB, CPU RSS=%.2f MB",544            gpu_after_warmup / (1024**2),545            cpu_after_warmup,546        )547 548    # GPU memory should not grow beyond warmup baseline.549    assert gpu_final <= gpu_after_warmup, (550        f"GPU memory leak detected! Warmup: {gpu_after_warmup / 1024**2:.2f} MB, "551        f"Final: {gpu_final / 1024**2:.2f} MB")552 553    # CPU RSS should not grow more than 50 MB over warmup (allows for minor554    # Python/CUDA runtime overhead but catches real leaks).555    cpu_growth = cpu_final - cpu_after_warmup556    assert cpu_growth < 50, (557        f"CPU memory leak detected! Growth: {cpu_growth:.2f} MB over "558        f"{num_steps - 2} steps (warmup={cpu_after_warmup:.2f} MB, "559        f"final={cpu_final:.2f} MB)")560 561    set_ns_compile(True)562    if rank == 0:563        logger.info("PASSED: test_leak (GPU stable, CPU growth=%.2f MB)",564                    cpu_growth)565 566 567def test_state_dict_save_load(rank, world_size):568    """Verify state_dict() works after offload and load_state_dict() resumes correctly.569 570    Uses torch.distributed.checkpoint (DCP) for serialization, matching571    the actual LLM training checkpoint flow.  DCP handles DTensors natively572    so the roundtrip is bitwise exact.573    """574    import shutil575    import tempfile576 577    import torch.distributed.checkpoint as dcp578    from optimizer.muon import Muon579    from optimizer.newton_schulz import set_ns_compile580 581    set_ns_compile(False)582    torch.manual_seed(42)583 584    mesh = _make_mesh(world_size)585 586    dim0, dim1 = 64, 128587    num_muon = 4588    num_adamw = 3589    num_steps = 3590 591    # Pre-generate all data.592    muon_init = [593        torch.randn(dim0, dim1, device="cuda") for _ in range(num_muon)594    ]595    adamw_init = [torch.randn(dim1, device="cuda") for _ in range(num_adamw)]596    all_grads_muon = [[597        torch.randn(dim0, dim1, device="cuda") for _ in range(num_muon)598    ] for _ in range(num_steps * 2)]599    all_grads_adamw = [[600        torch.randn(dim1, device="cuda") for _ in range(num_adamw)601    ] for _ in range(num_steps * 2)]602 603    def make_optimizer(cpu_offload):604        mp = [605            torch.nn.Parameter(606                distribute_tensor(muon_init[i].clone(), mesh, [Shard(0)]))607            for i in range(num_muon)608        ]609        ap = [610            torch.nn.Parameter(611                distribute_tensor(adamw_init[i].clone(), mesh, [Shard(0)]))612            for i in range(num_adamw)613        ]614        param_groups = [615            {616                "params": mp,617                "names": [f"layer.{i}.weight" for i in range(num_muon)],618                "use_muon": True,619                "lr": 0.02,620                "weight_decay": 0.01,621                "momentum": 0.95,622                "nesterov": True,623                "ns_steps": 5,624                "none_grad": False,625                "adamw_betas": (0.9, 0.95),626                "adamw_eps": 1e-8,627            },628            {629                "params": ap,630                "use_muon": False,631                "lr": 1e-3,632                "weight_decay": 0.01,633                "adamw_betas": (0.9, 0.95),634                "adamw_eps": 1e-8,635            },636        ]637        optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)638        if cpu_offload:639            optim.turn_on_cpu_offload()640        return optim, mp, ap641 642    # --- Run one optimizer for first half, save state, then create TWO643    # fresh optimizers: ref loads via deepcopy, resumed loads via DCP.644    # Both are fresh โ†’ same internal cache state โ†’ isolates DCP fidelity.645    optim_off, mp_off, ap_off = make_optimizer(True)646 647    for step_idx in range(num_steps):648        for i in range(num_muon):649            mp_off[i].grad = distribute_tensor(650                all_grads_muon[step_idx][i].clone(), mesh, [Shard(0)])651        for i in range(num_adamw):652            ap_off[i].grad = distribute_tensor(653                all_grads_adamw[step_idx][i].clone(), mesh, [Shard(0)])654        optim_off.step()655 656    with pytest.raises(657            RuntimeError,658            match="turn_off_cpu_offload\\(\\) before checkpoint save"):659        optim_off.state_dict()660 661    optim_off.turn_off_cpu_offload()662    sd_off = optim_off.state_dict()663 664    # Verify state tensors are NOT empty in the state_dict.665    for param_states in sd_off["state"].values():666        for key, val in param_states.items():667            if isinstance(val, torch.Tensor) and val.is_floating_point():668                assert val.untyped_storage().size() > 0, (669                    f"state_dict() returned empty storage for key '{key}' โ€” "670                    f"offload reload is broken")671 672    if rank == 0:673        logger.info("state_dict() contains valid (non-empty) tensors")674 675    # Save state tensors via DCP (matches real LLM training checkpoint flow).676    # Flatten state tensors with string keys for DCP compatibility.677    flat_state = {}678    for param_idx, param_state in sd_off["state"].items():679        for key, val in param_state.items():680            if isinstance(val, torch.Tensor):681                flat_state[f"state.{param_idx}.{key}"] = val682 683    # All ranks must use the same checkpoint directory.684    if rank == 0:685        ckpt_dir = tempfile.mkdtemp(prefix="cpu_offload_test_")686    else:687        ckpt_dir = ""688    ckpt_dir_list = [ckpt_dir]689    dist.broadcast_object_list(ckpt_dir_list, src=0)690    ckpt_dir = ckpt_dir_list[0]691    try:692        dcp.save(flat_state, checkpoint_id=ckpt_dir)693        dist.barrier()694 695        if rank == 0:696            logger.info("DCP save completed to %s", ckpt_dir)697 698        # --- Reference: fresh optimizer, load via deepcopy (no serialization).699        optim_ref, mp_ref, ap_ref = make_optimizer(True)700        for i in range(num_muon):701            mp_ref[i].data = mp_off[i].data.clone()702        for i in range(num_adamw):703            ap_ref[i].data = ap_off[i].data.clone()704        with pytest.raises(705                RuntimeError,706                match="turn_off_cpu_offload\\(\\) before checkpoint load"):707            optim_ref.load_state_dict(copy.deepcopy(sd_off))708        optim_ref.turn_off_cpu_offload()709        optim_ref.load_state_dict(copy.deepcopy(sd_off))710        optim_ref.turn_on_cpu_offload()711 712        # --- Resumed: fresh optimizer, load via DCP.713        optim_resumed, mp_resumed, ap_resumed = make_optimizer(True)714        for i in range(num_muon):715            mp_resumed[i].data = mp_off[i].data.clone()716        for i in range(num_adamw):717            ap_resumed[i].data = ap_off[i].data.clone()718 719        flat_target = {k: torch.zeros_like(v) for k, v in flat_state.items()}720        dcp.load(flat_target, checkpoint_id=ckpt_dir)721        dist.barrier()722 723        sd_loaded = copy.deepcopy(sd_off)724        for param_idx, param_state in sd_loaded["state"].items():725            for key in list(param_state.keys()):726                flat_key = f"state.{param_idx}.{key}"727                if flat_key in flat_target:728                    param_state[key] = flat_target[flat_key]729        with pytest.raises(730                RuntimeError,731                match="turn_off_cpu_offload\\(\\) before checkpoint load"):732            optim_resumed.load_state_dict(copy.deepcopy(sd_loaded))733        optim_resumed.turn_off_cpu_offload()734        optim_resumed.load_state_dict(sd_loaded)735        optim_resumed.turn_on_cpu_offload()736 737        if rank == 0:738            logger.info("Both optimizers loaded, starting comparison steps")739 740    finally:741        dist.barrier()742        if rank == 0:743            shutil.rmtree(ckpt_dir, ignore_errors=True)744 745    # Second half: reference continues, resumed uses loaded state.746    for step_idx in range(num_steps, num_steps * 2):747        for i in range(num_muon):748            g = all_grads_muon[step_idx][i]749            mp_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])750            mp_resumed[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])751        for i in range(num_adamw):752            g = all_grads_adamw[step_idx][i]753            ap_ref[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])754            ap_resumed[i].grad = distribute_tensor(g.clone(), mesh, [Shard(0)])755        optim_ref.step()756        optim_resumed.step()757 758    # Compare final params: bitwise exact (DCP preserves DTensor identity).759    for i in range(num_muon):760        ref_full = mp_ref[i].data.full_tensor()761        res_full = mp_resumed[i].data.full_tensor()762        torch.testing.assert_close(ref_full, res_full, atol=0, rtol=0)763 764    for i in range(num_adamw):765        ref_full = ap_ref[i].data.full_tensor()766        res_full = ap_resumed[i].data.full_tensor()767        torch.testing.assert_close(ref_full, res_full, atol=0, rtol=0)768 769    # Verify offload is active on the resumed optimizer.770    for p in mp_resumed:771        state = optim_resumed.state[p]772        if "momentum_buffer" in state:773            buf = state["momentum_buffer"]774            local_buf = buf._local_tensor if isinstance(buf, DTensor) else buf775            assert local_buf.untyped_storage().size() == 0, (776                "Resumed optimizer should have offloaded state after step()")777 778    set_ns_compile(True)779    if rank == 0:780        logger.info("PASSED: test_state_dict_save_load")781 782 783def test_checkpoint_memory(rank, world_size):784    """Verify checkpoint APIs require offload to be disabled explicitly."""785    from optimizer.muon import Muon786    from optimizer.newton_schulz import set_ns_compile787 788    set_ns_compile(False)789    torch.manual_seed(42)790 791    mesh = _make_mesh(world_size)792 793    dim0, dim1 = 512, 1024794    num_params = 8795 796    params, names = [], []797    for i in range(num_params):798        full = torch.randn(dim0, dim1, device="cuda")799        dt = distribute_tensor(full, mesh, [Shard(0)])800        p = torch.nn.Parameter(dt)801        p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),802                                   mesh, [Shard(0)])803        params.append(p)804        names.append(f"layer.{i}.weight")805 806    param_groups = [{807        "params": params,808        "names": names,809        "use_muon": True,810        "lr": 0.02,811        "weight_decay": 0.01,812        "momentum": 0.95,813        "nesterov": True,814        "ns_steps": 5,815        "none_grad": False,816    }]817    optim = Muon(params=param_groups, chunk_size=2, warmup_step=1)818    optim.turn_on_cpu_offload()819 820    # Step 1: run a step so offload initializes.821    optim.step()822    torch.cuda.synchronize()823 824    mem_after_step = torch.cuda.memory_allocated()825 826    # Calculate expected state size (momentum buffers, bf16).827    state_bytes = 0828    for p in params:829        state = optim.state[p]830        if "momentum_buffer" in state:831            buf = state["momentum_buffer"]832            local = buf._local_tensor if isinstance(buf, DTensor) else buf833            # Storage is freed, so use the tracked size.834            state_bytes += optim._cpu_offload_pool._storage_nbytes[id(buf)]835 836    if rank == 0:837        logger.info(838            "After step (offloaded): GPU alloc=%.2f MB, expected state size=%.2f MB",839            mem_after_step / 1024**2,840            state_bytes / 1024**2,841        )842 843    with pytest.raises(844            RuntimeError,845            match="turn_off_cpu_offload\\(\\) before checkpoint save"):846        optim.state_dict()847 848    optim.turn_off_cpu_offload()849    torch.cuda.synchronize()850    mem_after_turn_off = torch.cuda.memory_allocated()851    sd_for_load = copy.deepcopy(optim.state_dict())852 853    if rank == 0:854        logger.info(855            "After turn_off_cpu_offload: GPU alloc=%.2f MB",856            mem_after_turn_off / 1024**2,857        )858 859    assert mem_after_turn_off > mem_after_step, (860        f"turn_off_cpu_offload() should reload states to GPU. "861        f"After offload: {mem_after_step / 1024**2:.2f} MB, "862        f"After turn_off: {mem_after_turn_off / 1024**2:.2f} MB")863 864    optim.turn_on_cpu_offload()865    torch.cuda.synchronize()866    mem_after_turn_on = torch.cuda.memory_allocated()867 868    if rank == 0:869        logger.info("After turn_on_cpu_offload: GPU alloc=%.2f MB",870                    mem_after_turn_on / 1024**2)871 872    assert mem_after_turn_on <= mem_after_step + 4 * 1024 * 1024, (873        f"turn_on_cpu_offload() should return memory to offloaded level. "874        f"Expected <= {mem_after_step / 1024**2:.2f} MB (+4 MB tolerance), "875        f"got {mem_after_turn_on / 1024**2:.2f} MB")876 877    for p in params:878        p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),879                                   mesh, [Shard(0)])880    optim.step()881    torch.cuda.synchronize()882 883    mem_after_next_step = torch.cuda.memory_allocated()884 885    if rank == 0:886        logger.info(887            "After next step (re-offloaded): GPU alloc=%.2f MB",888            mem_after_next_step / 1024**2,889        )890 891    # Allow 4 MB tolerance for CUDA allocator fragmentation.892    assert mem_after_next_step <= mem_after_step + 4 * 1024 * 1024, (893        f"Memory should return to offloaded level after step(). "894        f"Expected <= {mem_after_step / 1024**2:.2f} MB (+4 MB tolerance), "895        f"got {mem_after_next_step / 1024**2:.2f} MB")896 897    with pytest.raises(898            RuntimeError,899            match="turn_off_cpu_offload\\(\\) before checkpoint load"):900        optim.load_state_dict(copy.deepcopy(sd_for_load))901 902    optim.turn_off_cpu_offload()903    optim.load_state_dict(sd_for_load)904    torch.cuda.synchronize()905 906    mem_after_load = torch.cuda.memory_allocated()907 908    if rank == 0:909        logger.info(910            "After load_state_dict with offload disabled: GPU alloc=%.2f MB",911            mem_after_load / 1024**2,912        )913 914    assert mem_after_load >= mem_after_turn_off, (915        "Loaded optimizer state should stay on GPU while offload is disabled")916 917    optim.turn_on_cpu_offload()918    torch.cuda.synchronize()919 920    pool = optim._cpu_offload_pool921    assert pool._initialized, (922        "Offload pool should be initialized after re-enabling offload")923    for grp in pool._groups.values():924        assert grp["cpu_flat"].is_pinned(), "CPU buffer must be pinned"925 926    # Step 5: verify the loaded optimizer can still step correctly.927    for p in params:928        p.grad = distribute_tensor(torch.randn(dim0, dim1, device="cuda"),929                                   mesh, [Shard(0)])930    optim.step()931    torch.cuda.synchronize()932 933    mem_final = torch.cuda.memory_allocated()934    assert mem_final <= mem_after_step + 4 * 1024 * 1024, (935        f"Final memory should be at offloaded level. "936        f"Expected <= {mem_after_step / 1024**2:.2f} MB (+4 MB tolerance), "937        f"got {mem_final / 1024**2:.2f} MB")938 939    set_ns_compile(True)940    if rank == 0:941        logger.info("PASSED: test_checkpoint_memory")942 943 944def main():945    rank, world_size = _setup()946 947    try:948        test_correctness(rank, world_size)949        test_memory(rank, world_size)950        test_adamw_offload(rank, world_size)951        test_memory_savings(rank, world_size)952        test_toggle_correctness(rank, world_size)953        test_leak(rank, world_size)954        test_state_dict_save_load(rank, world_size)955        test_checkpoint_memory(rank, world_size)956 957        if rank == 0:958            logger.info("=" * 50)959            logger.info("ALL CPU OFFLOAD TESTS PASSED")960            logger.info("=" * 50)961    finally:962        dist.destroy_process_group()963 964 965if __name__ == "__main__":966    main()967