Skip to content

Fix ZeroFlow sparse RNG to follow parameter device - #16

Open
kabishou11 wants to merge 1 commit into
THUDM:mainfrom
kabishou11:fix/zeroflow-device-aware-rng
Open

kabishou11 wants to merge 1 commit into
THUDM:mainfrom
kabishou11:fix/zeroflow-device-aware-rng

Conversation

@kabishou11

Copy link
Copy Markdown

Summary

Fixes #8.

ZeroFlow previously built sparse_grad_rng with cuda if torch.cuda.is_available() else cpu, ignoring the actual parameter device. On CPU/MPS models that still see CUDA available, sparse masking can then pair a CUDA generator with non-CUDA tensors.

This change:

  • Builds the generator from next(model.parameters()).device (cpu for non-cuda devices, since torch.Generator only supports cpu/cuda)
  • Removes the duplicate SimpleNamespace import
  • Removes the redundant get_grad_reduce call in SAM (already done in InftyBaseOptimizer.__init__)

Test plan

  • uv run --with pytest pytest -q tests/optim/test_zeroth_order_updates.py — 3 passed

Construct sparse_grad_rng from the model parameter device instead of
hardcoding cuda-if-available, so CPU/MPS models do not pair CUDA
generators with non-CUDA tensors during sparse masking. Also drop the
duplicate SimpleNamespace import and the redundant SAM get_grad_reduce
call already performed by InftyBaseOptimizer.

Fixes THUDM#8

This branch has not been deployed

No deployments
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.

Missing device-aware RNG construction in ZeroFlow

1 participant