Repository navigation
Conversation
Research-only writeup for RL-Align#386 (claimed by haoruilee). Maps closest GEMM/linear templates on upstream main, pins the FP32 API contract, lists files to add, and records test shapes plus open questions. No kernels. Co-authored-by: 0x4C33 <haoruilee@users.noreply.github.com>
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.
Research-only brief for RL-Align/RL-Kernel#386 (
latent_proj_gemm, claimed by @haoruilee). No kernels.Context was synced from
upstream/main@ccb70e3(this fork is 596 commits behind). Implementers should branch from upstream, not from this fork tip.1. Closest ops / layout to copy
det_gemm(PR [WS1][kernels] Batch-invariant deterministic GEMM (fwd + bwd) RL-Align/RL-Kernel#180) — mid-split K-tree, no Split-K, CUDA+Triton+gtest+bench+docs.det_gemmas-is: BF16-onlycheck_in, BF16 internal tree adds, unbounddet_gemm_fwd_fp32.lm_head— HF[out,in], optional bias,forward/forward_fp32.2. API contract
One registry op, two pinned directions, FP32 throughout:
xweight(HF)biasimg_in[..., 64][3072, 64][3072](checkpoint has it)proj_out[..., 3072][64, 3072][64]Y = X @ W.T + B, mid-split K-tree leaf 32, FP32 leaves and FP32 combines, no BF16 store. CUDA is bit-reference; Triton must match byte-for-byte; no silent fallback; WS2 = no.Official Diffusers +
Qwen/Qwen-Imagecheckpoint both shipimg_in.bias/proj_out.biaseven though the issue table omits bias.3. Files to add/change
New:
ops/{pytorch,triton,cuda}/linear/latent_proj_gemm.py,csrc/cuda/gemm/latent_proj_gemm_kernel.cu,tests/test_latent_proj_gemm.py,benchmarks/benchmark_latent_proj_gemm.py,docs/operators/latent-proj-gemm.md.Edit:
csrc/ops.cpp,setup.py(always-oncuda_sources, not SM90-gated),_C.pyi,registry.py,gtest/operator_{specs,inputs}.py, nav/README,__init__.pyexports.4. Test shapes / harness
Packed token counts after VAE
/8+ 2×2 pack:MPlus
M ∈ {1,7,31,64,128,256}. Gold = same mid-split FP32 tree (nottorch.matmul).check_operator.py --op latent_proj_gemm --dtype fp32after OP_SPECS lands.5. Open questions for RL-Align#386
Full writeup:
docs/design/qwen-image-ws1-latent-proj-gemm-brief.md.