Motif-Technologies/optimizer
58235
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 