Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion tests/unit/alexnet_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,7 +128,11 @@ def train_cifar(model, config, num_steps=400, average_dp_losses=True, fp16=True,
fork_kwargs = {"device_type": get_accelerator().device_name()}
else:
fork_kwargs = {}
with get_accelerator().random().fork_rng(devices=[get_accelerator().current_device_name()], **fork_kwargs):
# fork_rng only needs entries for backends with per-device generators: the global

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Add the required commit sign-off

This is a non-merge commit, but its message has no Signed-off-by trailer. Recreate the commit with --signoff so it satisfies the repository's commit requirements.

AGENTS.md reference: AGENTS.md:L8-L8

Useful? React with 👍 / 👎.

# CPU RNG is always saved, and torch.cpu has no get_rng_state to call anyway.
device_mod = torch.get_device_module(get_accelerator().device_name())

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Guard the device-module lookup on older PyTorch

When these tests run with the supported minimum PyTorch 2.0 (requirements/requirements.txt specifies torch>=2.0.0), torch.get_device_module does not exist, so every train_cifar call now raises AttributeError before entering fork_rng, including CUDA paths that previously worked. Please use a lookup available on older supported releases or guard this API by PyTorch version.

Useful? React with 👍 / 👎.

fork_devices = [get_accelerator().current_device_name()] if hasattr(device_mod, 'get_rng_state') else []
with get_accelerator().random().fork_rng(devices=fork_devices, **fork_kwargs):
ds_utils.set_random_seed(seed)

# disable dropout
Expand Down
Loading