Skip to content

fix: _patch_cpu_offload_apply handles .to(cuda) in addition to .cuda() - #59

Closed
cennn wants to merge 4 commits into
mainfrom
fix/cpu-offload-to-compat
Closed

fix: _patch_cpu_offload_apply handles .to(cuda) in addition to .cuda()#59
cennn wants to merge 4 commits into
mainfrom
fix/cpu-offload-to-compat

Conversation

@cennn

@cennn cennn commented Aug 18, 2026

Copy link
Copy Markdown
Collaborator

Summary

  • _patch_cpu_offload_apply only recognized .cuda() lambda but not .to("cuda:X") lambda. When callers use model.to(device) instead of model.cuda(), the offload hook fell through to _orig_apply, putting all weights on GPU and defeating model_cpu_offload.
  • Fix: probe the .to() lambda with a small CPU tensor to detect if the target device is CUDA, and treat it the same as .cuda() for offload interception.

Test plan

  • Existing test_cpu_offload_placement passes (uses .cuda())
  • New scenario: model.to("cuda:0") with offload enabled keeps decorated module weights on CPU
  • model.to("cpu") does not trigger _force_cpu (probe returns is_cuda=False)

cennn added 4 commits August 18, 2026 11:48
…uda()

When callers use model.to(cuda:X) instead of model.cuda(), the
_cpu_apply hook did not recognize the Module.to lambda and fell through
to _orig_apply, putting all weights on GPU. This defeats the purpose of
model_cpu_offload.

Fix: probe the .to() lambda with a small CPU tensor to detect if the
target device is CUDA, and treat it the same as .cuda() for offload
interception.
The scheduler hard-coded two layers of lookahead. Prefetching costs GPU memory
that a small card may not have, so expose it as max_prefetch_lookahead and keep
2 as the default; 0 disables prefetch entirely.
max_tiles=3 can overflow CUDA z-grid limit (65535) on conv-heavy
dynamic-shape graphs when coalesce_tiling_analysis collapses high-dim
tensors into 3D. Grid3D has no z-overflow handling unlike
Grid2DWithYZOverflow. The previous PT-version guard (2 on PT>=2.12,
3 otherwise) left PT<2.12 exposed. Cap at 2 on all versions since
max_tiles=3 is documented as experimental.
The pass hardcoded max_tiles=2 after we found Grid3D has no z-overflow
handling. That value is right as a default but there was no way to try 3
without editing the pass. Expose it as PassConfig.nd_tiling_max_tiles,
reachable via MAGI_COMPILE_PASS_CONFIG__ND_TILING_MAX_TILES.

Turning the whole pass off is not an alternative: it also sets
prefer_nd_tiling and tile_reductions, which conv-heavy dynamic-shape
graphs need to avoid degrading to Grid1D.
@cennn

cennn commented Aug 21, 2026

Copy link
Copy Markdown
Collaborator Author

Superseded by #61.

Same 4-file / +33 -8 diff vs main. #61 is a clean cherry-pick of these commits onto current main (single commit, title/body cover all four changes). Closing this as a duplicate.

@cennn cennn closed this Aug 21, 2026
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.

1 participant