fix(dmd): only pass driver="gesvd" to torch.linalg.svd on CUDA - #1116
Merged
Merged
Conversation
`dmd_svd_precision="high"` called torch.linalg.svd(..., driver="gesvd") on every device. PyTorch rejects `driver=` for non-CUDA inputs, and _dmd_fit_one swallows the error and returns None, so on CPU the DMD calibrator silently fell back to reusing the last snapshot. Mirror svdquant/lowrank.py: use gesvd on CUDA, the default driver elsewhere. Add a CPU test covering all three precision levels.
Contributor
There was a problem hiding this comment.
Copilot review overview
🟢 Approval recommended
The focused fix follows an existing repository pattern and is covered by an appropriate regression test.
Review effort: Balanced
Findings: None
What changed in this PR
Fixes CPU/non-CUDA DMD forecasting when using high SVD precision.
Changes:
- Applies the
gesvddriver only to CUDA tensors. - Adds CPU regression coverage for every precision level.
| File | Description |
|---|---|
src/cache_dit/caching/cache_contexts/calibrators/dmd.py |
Guards the CUDA-only SVD driver. |
tests/api/test_dmd_svd_precision.py |
Verifies accurate CPU forecasts. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
With
dmd_svd_precision="high",_dmd_svd()callstorch.linalg.svd(X, full_matrices=False, driver="gesvd")on every device. PyTorch only acceptsdriver=for CUDA inputs:_dmd_fit_one()catches that exception as a degenerate fit and returnsNone, so on CPU (and any non-CUDA device) the DMD calibrator never forecasts with"high". It quietly reuses the last snapshot, with no warning, while"low"and"medium"work.Repro on CPU (rank-3 linear dynamics, which DMD represents exactly; 6 snapshots, forecast step 7):
Change
dmd.py: usedriver="gesvd"only whenX.device.type == "cuda", and the default driver otherwise. This is the same guardsvdquant/lowrank.pyalready uses (driver="gesvd" if weight.device.type == "cuda" else None). The_dmd_svddocstring now says"high"is CUDA-only and matches"medium"elsewhere.tests/api/test_dmd_svd_precision.py: a CPU test that runsDMDStatefor all three precision levels on a rank-3 trajectory and checks that the forecast matches step 7 and is not just the newest snapshot.CUDA behaviour is unchanged.
Validation (CPU only, no GPU available)
Python 3.11, torch 2.14.0+cpu, diffusers 0.40.0, on
mainat 99fc5a8. Rebased onto 5c2d065 (only README changed since) and reran the new test:3 passed.main:1 failed, 2 passed([high]fails). With the fix:3 passed.pytest tests/api tests/utilswith the fix:112 passed, 25 skipped, 1 failed. The one failure istest_load_configs[api/config.yaml], which uses a path relative totests/; it passes when run fromtests/(1 passed) and is unrelated to this change.flake8 --config=setup.cfgis clean on both files; the test file is formatted with the repo's.style.yapf.This patch was prepared with the help of an AI coding assistant (Claude Code).