Skip to content

fix(dmd): only pass driver="gesvd" to torch.linalg.svd on CUDA - #1116

Merged
DefTruth merged 1 commit into
vipshop:mainfrom
kv-248:fix-dmd-svd-high-cpu
Sep 29, 2026
Merged

DefTruth merged 1 commit into
vipshop:mainfrom
kv-248:fix-dmd-svd-high-cpu

Conversation

@kv-248

@kv-248 kv-248 commented Sep 28, 2026

Copy link
Copy Markdown
Contributor

Problem

With dmd_svd_precision="high", _dmd_svd() calls torch.linalg.svd(X, full_matrices=False, driver="gesvd") on every device. PyTorch only accepts driver= for CUDA inputs:

RuntimeError: torch.linalg.svd: keyword argument `driver=` is only supported on CUDA inputs with cuSOLVER backend.

_dmd_fit_one() catches that exception as a degenerate fit and returns None, 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):

low    rel_err=7.32e-08 returned_newest_snapshot=False
medium rel_err=7.32e-08 returned_newest_snapshot=False
high   rel_err=1.02e-01 returned_newest_snapshot=True

Change

  • dmd.py: use driver="gesvd" only when X.device.type == "cuda", and the default driver otherwise. This is the same guard svdquant/lowrank.py already uses (driver="gesvd" if weight.device.type == "cuda" else None). The _dmd_svd docstring now says "high" is CUDA-only and matches "medium" elsewhere.
  • tests/api/test_dmd_svd_precision.py: a CPU test that runs DMDState for 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 main at 99fc5a8. Rebased onto 5c2d065 (only README changed since) and reran the new test: 3 passed.

  • New test on unpatched main: 1 failed, 2 passed ([high] fails). With the fix: 3 passed.
  • pytest tests/api tests/utils with the fix: 112 passed, 25 skipped, 1 failed. The one failure is test_load_configs[api/config.yaml], which uses a path relative to tests/; it passes when run from tests/ (1 passed) and is unrelated to this change.
  • flake8 --config=setup.cfg is clean on both files; the test file is formatted with the repo's .style.yapf.
  • Not tested: the CUDA path (no GPU here). The CUDA branch is unchanged code.

This patch was prepared with the help of an AI coding assistant (Claude Code).

`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.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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 gesvd driver 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.

@DefTruth DefTruth left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

LGTM~ Thx for this fix!

@DefTruth
DefTruth merged commit a7898aa into vipshop:main Sep 29, 2026
4 checks passed
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.

3 participants