kohya-ss / kohya-ss/sd-scripts

`torch.cuda.amp` in `__main__` self-test blocks and the IPEX hijack — deprecated since torch 2.3/2.4, unguarded on the main install path

Open
#2,420 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
7.2k
Forks
1.2k
Avg merge
11m
Merged PRs (30d)
2

Description

## Summary

17 references to the deprecated `torch.cuda.amp` API remain at HEAD (`37a1cbbc`). `torch.cuda.amp.GradScaler` has been deprecated since torch 2.3 and `torch.cuda.amp.autocast` since torch 2.4 — both still work but emit a `FutureWarning`, and the deprecation notice says they "will be removed in a future release". When that removal lands, every site below fails at import/run time.

Two groups:

**1. `__main__` self-test blocks — 8 references in 4 files, no version guard**

| File | Lines | Usage |
|---|---|---|
| `library/sdxl_original_unet.py` | 1266, 1282 | `torch.cuda.amp.GradScaler(enabled=True)` / `with torch.cuda.amp.autocast(enabled=True):` |
| `library/sdxl_original_control_net.py` | 236, 253 | `GradScaler(...)` / `with torch.cuda.amp.autocast(enabled=True, dtype=torch.bfloat16):` |
| `networks/control_net_lllite.py` | 418, 434 | `GradScaler(...)` / `with torch.cuda.amp.autocast(enabled=True):` |
| `networks/control_net_lllite_for_train.py` | 471, 487 | `GradScaler(...)` / `with torch.cuda.amp.autocast(enabled=True, dtype=torch.bfloat16):` |

All 8 are inside `if __name__ == "__main__":` blocks (at `sdxl_original_unet.py:1241`, `sdxl_original_control_net.py:199`, `control_net_lllite.py:374`, `control_net_lllite_for_train.py:414`). None of the 4 files has any version guard (`LooseVersion` / `torch.__version__` / `hasattr(torch.amp, ...)` / try-except) around them. These are developer-facing smoke tests — the README and docs never mention running them — but they are the documented way to sanity-check those two UNet / ControlNet implementations.

**2. `library/ipex/hijacks.py` — 9 references in the IPEX monkey-patch, hasattr-guarded**

Lines 448-461 of `ipex_hijacks()`: `custom_fwd`/`custom_bwd` aliasing and `torch.cuda.amp = torch.xpu.amp` / `torch.cuda.amp = torch.amp` module swaps, guarded by `hasattr(torch.xpu, "amp")`, `hasattr(torch.amp, "custom_fwd")`, etc. This runs only on the Intel GPU (IPEX) initialization path — the guarded aliasing actually *replaces* `torch.cuda.amp` with the `torch.amp` implementation, so the deprecation concern here is lower.

## No torch version constraint

- `requirements.txt` has no torch entry; torch is pulled in only indirectly via `diffusers[torch]==0.32.1` (requirements.txt:3). `setup.py` has no `install_requires`.
- README.md:211 states this explicitly: *"The file does not contain requirements for PyTorch. Because the version of PyTorch depends on the environment, it is not included in the file."*
- README.md:213: *"The scripts are tested with PyTorch 2.6.0. PyTorch 2.6.0 or later is required."* and README.md:215 recommends PyTorch 2.8.0 for RTX 50-series — both of these already deprecate `torch.cuda.amp.GradScaler`/`autocast` (since 2.3/2.4), so anyone on the documented install path gets the `FutureWarning` when running one of the self-test blocks, and would break outright once torch removes the module.

## Suggested fix

Low-risk, 8 lines: swap the self-test blocks to the device-agnostic APIs.

```python
scaler = torch.amp.GradScaler("cuda", enabled=True)
with torch.amp.autocast("cuda", enabled=True, dtype=torch.bfloat16):
```

`torch.amp.GradScaler` exists since torch 2.3, `torch.amp.autocast` since torch 2.0, so this works on the torch ≥ 2.6.0 floor the README already declares — no version-branching needed. The `ipex/hijacks.py` aliasing is arguably fine as-is (it already targets `torch.amp`), but worth a comment.

Low priority: the main training path (`train_network.py`, `sdxl_train.py`, `train_util.py`) has zero `torch.cuda.amp` references; this only affects the undocumented self-test blocks and the IPEX path.

## Reference

- [PyTorch AMP docs](https://pytorch.org/docs/stable/amp.html) — official deprecation notice

- [huggingface/lerobot#3167](https://github.com/huggingface/lerobot/pull/3167) — same migration in LeRobot, merged
- [#1649](https://github.com/kohya-ss/sd-scripts/issues/1649) — closed; the only amp-related issue, and it is about a `FutureWarning` coming from torch's own `torch.utils.checkpoint.py`, not this repo's code

Contributor guide

No contributing guide indexed for this repository

Research direction

Read the __main__ self-test blocks in library/sdxl_original_unet.py, library/sdxl_original_control_net.py, networks/control_net_lllite.py, and networks/control_net_lllite_for_train.py, then inspect the guarded aliases in library/ipex/hijacks.py. Update the eight self-test references to the device-agnostic torch.amp APIs supported by the documented PyTorch floor, and confirm the self-test blocks no longer use deprecated torch.cuda.amp calls.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
testing-qa, tooling
Issue type
Bug
Difficulty
3/5
Estimated time
1-2 days
Activity status
Active
Clarity
Clearly specified
Newbie friendliness
78/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.