対象 GPU(確認済み): NVIDIA RTX 5090、Blackwell、compute capability sm_120、VRAM 32GB、driver 610.74。
現状(確認済み): torch 2.11.0+cpu と jax 0.7.1(CpuDevice)しか入っておらず、いま GPU では走らない。sm_120 は CUDA 12.8+ 世代のビルドが要る。
このカタログの性格: Web 検索は使わず、RAD コーパス(D:/docs/*_corpus_v2/、実在確認済み)+ 知識ベースで書いた。推測は「(推測)」と明示する。実コード(physarum_search.py / shapematch.py / afterman/eco_world.py)を読んで各ワークロードに紐づけた。
出典表記:
| # | ワークロード | 計算の核 | いまのボトルネック | バッチ次元 | 精度要件 |
|---|---|---|---|---|---|
| 1 | 粘菌ソルバ physarum_search.py |
ラプラシアン L(D)p=b の線形ソルブ × 時間反復(伝導率 D を毎反復更新) |
dense n×n を毎反復 torch.linalg.solve(O(n³))。しかも毎反復 A を torch.zeros(n,n) で作り直す |
多数の迷路 × 多数のパラメータ(μ, dt, D_init) | fp64(コード実測: torch.float64 固定。ラプラシアンは条件数が悪くなりやすい) |
| 2 | 形状マッチング shapematch.py |
テンプレートのエッジ勾配ベクトルと画像勾配の内積を、多数の位置 × スケール × 角度で評価 | Python 二重ループ _scan_flat / _score_at(fancy-index の内積)。scale/角度も Python ループ |
位置 × スケール × 角度 × 複数インスタンス | fp32 で十分(勾配方向の内積。相関のダイナミックレンジは狭い) |
| 3 | Afterman 進化 afterman/eco_world.py |
個体群の並列前進評価(RNN 政策) + 構造進化 | すでに JAX で綺麗にバッチ化済(個体ごと重み w_in/w_rec/w_out を einsum("nij,nj->ni")、lax.scan で T ステップ、jit(static_argnums=1)) |
個体数 n(集団) | fp32/bf16 で十分(進化 fitness は近似で足りる) |
重要な非対称性: 3 番はすでに GPU-ready の書き方(SoA + batched einsum + scan)。1・2 は「Python ループ + 逐次 dense ソルブ」で GPU 化の伸びしろが最大。着手優先度は 1・2 が上、3 は wheel を入れれば jax.devices() が cuda になるだけでほぼ動く(詳細は §5)。
各パターン = 名前 / いつ使う / 落とし穴 / 我々のどれに効くか。
if cfg.contested: のような分岐はバッチ全体で同じ枝にして warp 発散を避ける [知識]。torch.linalg.solve はバッチ版を持つ(先頭次元がバッチ)。(B, n, n) を渡せば B 個の系を1呼び出しで解く。cuSOLVER の batched API、cuBLAS の *gemmStridedBatched が裏で効く [知識]。疎なら torch-sla が「共有 or 個別の sparsity pattern に対する batched solve」を明示サポート [コーパス: torch-sla, arXiv 2601.13994]。_score_at は pts に沿った fancy index で散らばったアクセス。GPU 化するなら im2col でパッチを連続メモリに展開してから GEMM に落とす(P10)と coalesced になる [知識]。.cpu().numpy() / float(...) / .item() を呼ぶと毎回同期 + 転送が入り、GPU が待たされる。physarum_search.py:147 d = float(torch.max(torch.abs(newD - D))) を毎反復。これは毎反復 device→host 同期。GPU 化したら収束判定を K 反復ごとにまとめるか、d を device 上に貯めて最後にまとめて host へ [知識]。history.append(d) も同様。GPU では history を tensor に貯めて最後に一括で返す。for が独立反復なら、その反復を tensor の1軸にする。_scan_flat の for r0 / for c0 → 全位置を一気にスコア地図として計算(=相関/畳み込み、P10)。scale/角度の for もバッチ軸へ:テンプレートを角度・scale ごとに回転/拡縮したスタック (S*A, h, w) を作り、画像との相関を1バッチで [知識]。(B, n, n) にして batched solve(P2)。torch.nn.functional.conv2d が最速。勾配2成分(gy,gx)を入力2チャンネル、テンプレの (grad_y, grad_x) を重みにして内積=2ch の conv の和で表せる。O(N log N)。小テンプレートでは FFT のオーバーヘッド負け。_score_at)。conv で内積を出した後に、正規化・閾値・カウントを elementwise で当てる2段構成にする。eco_world.py がすでにこの定石。w_rec * mask(eco_world.py:202)で構造(結線)を固定形状 (n,H,H) に載せ、mask を 0/1 で切る。個体ごとに違う構造を固定形状 + mask で表す=進化で幅を変えても同じ jit カーネルで回る(memory project_afterman_structure_evolution の思想と一致)。[(x,y,e), ...] でなく x[], y[], e[]。st["pos"], st["energy"], st["alive"] が別配列)。1 の Graph も edges/length/coords が別配列で SoA 寄り。torch.linalg.solve/cholesky(バッチ対応)、または cuSOLVER batched。cupyx.scipy.sparse.linalg(cg, spsolve)。torch.compile で融合(P15)。torch.linalg.solve のバッチ版で足りたと気づく。必ずベースラインを測ってから(memory feedback_beat_the_null_before_claiming)。einsum("nij,nj->ni") は batched matvec = GEMM 族。2 の im2col+GEMM(P10)。jit + lax.scan(eco_world.py:291)がすでにこれ。scan 全体が1つの融合実行になり、ステップごとのカーネル起動を畳む。jit 境界は rollout 全体(1ステップごとに jit しない)が正解=既にそう書けている。torch.compile で包むか、行列組み立て(index_put_/index_add_)+ solve + 更新を1つの compiled region に。torch.compile は動的形状で再コンパイル多発。形状を固定して(P11)。JAX も static_argnums に可変値を入れると再 jit(3 は T を static にしていて OK、ただし T を変える実験では再 jit されると認識)。physarum_search.py:133 A = torch.zeros(n, n, ...) を毎反復。n=5000 の迷路で fp64 なら 5000²×8B = 200MB を毎反復確保。バッチ B 個なら ×B で即 OOM。疎(CSR/COO)で持つのが必須。torch-sla / cupy sparse へ(P13-2)。float64 固定。
torch.set_float32_matmul_precision("high") / torch.backends.cuda.matmul.allow_tf32=True で有効化 [知識]。2・3 の GEMM に効く。1 の fp64 ソルブには効かない(TF32 は fp32 経路)。feedback_benchmark_honest_disclosure)。A を明示的に組むのが高い/メモリを食う。CG/BiCGSTAB は A@x の作用さえあればよい。ラプラシアンなら A@x = degree*x - (隣接からの寄与) を疎に計算でき、dense A を作らない。L(D) が毎反復わずかに変化。前反復の解 p を次反復 CG の初期値にすると反復数が激減 [知識]。torch.linalg.solve(直接法)では効かない。反復法に切り替えて初めて効く。つまり「dense 直接 → 疎反復」への移行とセット。| 落とし穴 | 症状 | 対策 | 該当 |
|---|---|---|---|
| 小問題を1個ずつ GPU 投入 | CPU より遅い | batched-kernel(P2)でまとめる | 1・2 |
反復中の .item()/float()/.cpu() |
GPU が毎反復待つ | 収束判定を K 反復まとめ、history は device に貯める(P7) | 1 |
dense A を毎反復確保 |
OOM / 帯域律速 | 疎 + matrix-free(P16/P18) | 1 |
| 全 scale×角度×位置を同時 materialize | VRAM 爆発 | ピラミッドで段階化(P9) | 2 |
| 安易な fp32 化 | 収束せず/経路が変わる | 混合精度反復改良、fp64 基準と照合(P17) | 1 |
| 動的形状で再コンパイル | jit/compile が毎回走る | 形状固定 + padding/mask(P11/P15) | 2・3 |
| warp 発散する分岐 | 実効並列度が落ちる | バッチ全体で同じ枝(P1) | 全部 |
| Triton を書いてから既製で足りたと気づく | 労力の無駄 | ベースライン測定 → 段階的に登る(P13) | 全部 |
| やりたいこと | 第一候補 | 代替 | 出典 |
|---|---|---|---|
dense バッチ線形系 (B,n,n) |
torch.linalg.solve(バッチ対応) |
cuSOLVER batched, cupy.linalg |
[知識] |
| 疎バッチ線形系(共有/個別 sparsity) | torch-sla(cuDSS/CuPy/torch-iterative 自動ディスパッチ) | cupyx.scipy.sparse.linalg.cg/spsolve |
[コーパス: arXiv 2601.13994] |
| 疎 CG/BiCGSTAB + プリコンディショナ | cupyx.scipy.sparse.linalg |
torch-sla iterative backend | [知識/コーパス] |
| 三角ソルブ(プリコンディショナ適用)を速く | 既製で測る → 足りねばサブドメイン化 | 自作 Triton/CUDA | [コーパス: arXiv 2508.04917] |
| 相関/畳み込み(形状マッチ) | F.conv2d(小テンプレ) |
FFT(大テンプレ)、im2col+GEMM | [知識] |
| 集団の並列前進 + scan | JAX vmap/lax.scan/jit(3 は既実践) |
torch.vmap + torch.compile |
[知識] |
| 乱数(集団・確率イベント) | jax.random.split(3 は既実践、eco_world.py:244) |
torch.Generator per-stream |
[知識] |
| elementwise 融合が律速 | torch.compile / Triton |
CUDA Graph(固定反復) | [コーパス: mlops doc_0717/0521] |
| 混合精度反復改良 | fp32 内側 + fp64 補正(rescaling で fp16 も可) | — | [コーパス: arXiv 2602.14450] |
| TF32 有効化(fp32 GEMM 高速化) | torch.set_float32_matmul_precision("high") |
allow_tf32=True |
[知識] |
JAX の要点(3 向け) [知識]:
vmap = 集団をバッチ軸に自動ベクトル化。3 は個体ごと重みを明示バッチ(n 軸)にしていて実質同義。lax.scan = 時間ループを1つの融合カーネルに(Python for より圧倒的に速い)。3 は rollout で実践済。jit の境界 = rollout 全体を1回 jit(ステップごとに jit しない)。3 は @partial(jax.jit, static_argnums=1) で T を static にして正解。jax.random.split = 反復ごとに key を割る(3 は step 内で split 実践)。同じ key を使い回すと相関した乱数になる落とし穴を回避済。w_rec*mask)。進化で幅が変わっても同じ jit 済カーネルで回る。現状 torch 2.11.0+cpu / jax 0.7.1 CpuDevice。sm_120(Blackwell)は CUDA 12.8+ 世代が要る。以下は当たりであり、入れる前にリリースノートで sm_120 対応を確認すること。
PyTorch(推測ベース、要確認):
# 既存 CPU 版を外してから CUDA 12.8 wheel を入れる想定(バージョンは要確認)
py -3.11 -m pip uninstall torch
py -3.11 -m pip install torch --index-url https://download.pytorch.org/whl/cu128
# cu129 系が出ていればそちらの方が Blackwell 対応が新しい可能性(推測)
no kernel image is available for execution on the device で落ちる。後者なら sm_120 対応 wheel を待つ/nightly を使う。py -3.11 -c "import torch; print(torch.__version__, torch.cuda.is_available(), torch.cuda.get_device_capability())" → (12,0) 相当が出れば sm_120 認識。JAX(推測ベース、要確認):
py -3.11 -m pip install -U "jax[cuda12]"
jax[cuda12] は CUDA 12 系 + cuDNN を pip で引く。Blackwell 対応は jaxlib の CUDA 同梱バージョン依存。0.7.1 で sm_120 が通らなければ新しい jaxlib(cuda12 plugin)へ上げる。py -3.11 -c "import jax; print(jax.devices())" → CudaDevice が出れば OK。共通の落とし穴(知識):
+cpu サフィックスの torch が残ると cuda.is_available() が False)。第1手: 粘菌ソルバ(1)の「dense 直接 → 疎 + matrix-free 反復 + warm start」への書き換え。ただし CPU/numpy でまず疎化して正しさを固定してから GPU。
A を再確保、毎反復 host 同期)。構造的欠陥が3つ重なっている(P16/P7/P18)ので伸びしろ最大。L(D) が少しずつ変わる性質が warm start(P20)にぴったり。前ステップ解を初期値にした CG は反復数が激減する見込み。feedback_cpu_short_poc_before_gpu に従い、まず CPU で疎化 + 反復法 + warm start を入れて最短路収束が保たれることを確認してから device="cuda"。fp64→混合精度の是非も CPU 段階で残差を見て決める。第2手: 形状マッチング(2)の「Python 二重ループ → conv2d/im2col バッチ」化。
_scan_flat/_score_at の二重 for は本質的に相互相関(P10)。F.conv2d に落とせば1手で大幅高速化、しかもピラミッド構造(粗→精密)がそのままバッチ段階化(P9)になり VRAM 爆発を避けられる。第3手: Afterman(3)は wheel を入れるだけ。コード変更はほぼ不要。
lax.scan + jit + random.split + mask で可変構造)。jax[cuda12](または WSL2)を入れれば jax.devices() が cuda になりそのまま乗る見込み。最初の1手は「粘菌ソルバを疎 + matrix-free 反復 + warm start に作り替える(まず CPU で正しさ確定)」。 現状の dense O(n³)・毎反復再確保・毎反復同期という3重の構造欠陥を潰し、時間反復と warm start の相性・batched 疎 solve ライブラリ(torch-sla)の存在という追い風が全部そろっているため、GPU 化の投資回収が最も速い。
すべて D:/docs/numerical_methods_corpus_v2/ と D:/docs/mlops_corpus_v2/ 配下で grep により実在確認(ファイルパスは本文脚注のクラスタ)。
TF32/FP8/Blackwell sm_120 の具体・wheel バージョン・fp64 レートはコーパス外の知識ベース + 推測であり、実インストール時にリリースノート/実測で確認すること。 </content> </invoke>
上のカタログで挙げた「第1手」を実装し、実際に GPU(loco venv の torch 2.11.0+cu128、
CUDA 12.8、RTX 5090)で回した結果。packages(imgevolve ルート)の
physarum_search.py / tests/test_physarum_search.py。
rs.max() の device→host 同期を
cg_check_every で間引く(GPU では 1 反復ごとの同期が直列化要因)。| 規模 | CPU 逐次疎(基準) | GPU 最良(FP32+CUDA graph) | 倍率 | |—|—|—|—| | k=64 (n=4,096, B=16) | 27.9 s | 1.12 s | 24.9x | | k=128 (n=16,384, B=16) | 89.1 s | 1.09 s | 82.2x | | k=200 (n=40,000, B=16) | 230.7 s | 1.42 s | 162.3x |