Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
148 commits
Select commit Hold shift + click to select a range
0522865
[WS1][CUDA][MiniMax-H3] timestep_sinusoid_h3: FP32 timestep features …
fusheng-ji Oct 6, 2026
8c0d167
[WS1][MiniMax-H3] timestep_sinusoid_h3: B200 evidence report and figure
fusheng-ji Oct 7, 2026
65ef7f6
[WS1][CUDA][MiniMax-H3] timestep_mlp_fp32: deterministic FP32 timeste…
fusheng-ji Oct 7, 2026
6383a08
[WS1][MiniMax-H3] timestep_mlp_fp32: B200 evidence report and figure
fusheng-ji Oct 7, 2026
d06e120
[WS1][CUDA][MiniMax-H3] adaln_projection_3mod: three-modality AdaLN p…
fusheng-ji Oct 7, 2026
a8cd1bf
[WS1][MiniMax-H3] adaln_projection_3mod: B200 evidence report and figure
fusheng-ji Oct 7, 2026
fa551c2
[WS1][CUDA][MiniMax-H3] adaln_row_gather: fused six-way AdaLN row gather
fusheng-ji Oct 7, 2026
022c09a
[WS1][MiniMax-H3] adaln_row_gather: B200 evidence report, chain repla…
fusheng-ji Oct 7, 2026
f36aa60
[WS1][CUDA][MiniMax-H3] h3_rmsnorm: RMSNorm with fused AdaLN modulation
fusheng-ji Oct 7, 2026
6041a7a
[WS1][MiniMax-H3] h3_rmsnorm: B200 evidence report and figure
fusheng-ji Oct 7, 2026
a26b41a
[WS1][CUDA][MiniMax-H3] adaln_gate_residual: gated residual with in-k…
fusheng-ji Oct 7, 2026
e8e0d6e
[WS1][MiniMax-H3] adaln_gate_residual: B200 evidence report and figure
fusheng-ji Oct 7, 2026
38d575f
[WS1][CUDA][MiniMax-H3] final_adaln_out: norm_out projection + final …
fusheng-ji Oct 7, 2026
2a4774d
[WS1][MiniMax-H3] final_adaln_out: B200 evidence report, chain replay…
fusheng-ji Oct 7, 2026
eed949c
fix(h3): reject mixed CUDA devices in deterministic linear helpers
fusheng-ji Oct 7, 2026
f32ef1f
fix(h3): validate native timesteps and extracted weight contracts
fusheng-ji Oct 7, 2026
5121e19
Merge timestep sinusoid review fixes into H3 timestep MLP
fusheng-ji Oct 7, 2026
26ee42a
Merge inherited RFC 420 review fixes into PR #485
fusheng-ji Oct 7, 2026
5e5f284
Merge inherited RFC 420 review fixes into PR #486
fusheng-ji Oct 7, 2026
ee8c845
Merge inherited RFC 420 review fixes into PR #487
fusheng-ji Oct 7, 2026
6ba92e4
Merge inherited RFC 420 review fixes into PR #488
fusheng-ji Oct 7, 2026
a236ef4
Merge inherited RFC 420 review fixes into PR #489
fusheng-ji Oct 7, 2026
b62208d
fix(h3): address AdaLN projection review findings
fusheng-ji Oct 7, 2026
fedd7d7
fix(h3): validate native gather indices and replay stage dependencies
fusheng-ji Oct 7, 2026
082b9e4
ci: enable CodeRabbit auto review for test-h3
fusheng-ji Oct 7, 2026
b9ed910
ci: enable CodeRabbit automatic reviews for test-h3
fusheng-ji Oct 7, 2026
3189a74
ci: inherit existing CodeRabbit review settings
fusheng-ji Oct 7, 2026
394663d
docs(h3): document conditioning operators and review regression tests
fusheng-ji Oct 7, 2026
40542b5
test(h3): enforce complete tensor pairings in regression checks
fusheng-ji Oct 7, 2026
6af5bca
fix(h3): preserve gradients for empty linear dimensions
fusheng-ji Oct 7, 2026
80e4609
fix(h3): validate RMSNorm inputs and isolate backward benchmarks
fusheng-ji Oct 7, 2026
6a413a1
docs(h3): refresh B200 evidence with backward-only alternating samples
fusheng-ji Oct 7, 2026
12c0b7d
fix(h3): reject mixed-device RMSNorm native inputs
fusheng-ji Oct 7, 2026
372d6d7
fix(h3): retain FP64 final AdaLN golden gradients
fusheng-ji Oct 7, 2026
e565b45
fix(h3): address native validation and evidence review comments
fusheng-ji Oct 7, 2026
108ec76
docs(h3): regenerate final AdaLN evidence against FP64 gradients
fusheng-ji Oct 7, 2026
4947996
[WS2][CUDA][MiniMax-H3] tp_adaln_3mod: column-parallel AdaLN projecti…
fusheng-ji Oct 7, 2026
6c0d380
[WS2][CUDA][MiniMax-H3] sp_norm_adaln: sequence-parallel norm/AdaLN/g…
fusheng-ji Oct 7, 2026
b43ad26
[WS2][MiniMax-H3] tp_adaln_3mod: 8 x B200 NCCL evidence report and fi…
fusheng-ji Oct 7, 2026
d371c6d
Merge tp_adaln_3mod evidence into PR #494
fusheng-ji Oct 7, 2026
4992746
[WS2][MiniMax-H3] sp_norm_adaln: 8 x B200 NCCL evidence report and fi…
fusheng-ji Oct 7, 2026
001684d
fix(h3): store norm_out.norm and 1 + scale in BF16 in the final AdaLN…
fusheng-ji Oct 8, 2026
3f5d333
docs(h3): regenerate final AdaLN evidence with the BF16 norm_out golden
fusheng-ji Oct 8, 2026
3507b74
test(h3): assert adaln_row_gather batch invariance per logical row
fusheng-ji Oct 8, 2026
46154fc
fix(h3): store norm(x) and 1 + scale in BF16 in the modulated RMSNorm…
fusheng-ji Oct 8, 2026
cce7ce0
Merge inherited RFC 420 review fixes into PR #486
fusheng-ji Oct 8, 2026
4d015a7
Merge inherited RFC 420 review fixes into PR #487
fusheng-ji Oct 8, 2026
a1c597d
Merge inherited RFC 420 review fixes into PR #488
fusheng-ji Oct 8, 2026
90cbde0
Merge inherited RFC 420 review fixes into PR #489
fusheng-ji Oct 8, 2026
609d585
Merge inherited RFC 420 review fixes into PR #493
fusheng-ji Oct 8, 2026
707fd04
Merge inherited RFC 420 review fixes into PR #494
fusheng-ji Oct 8, 2026
59c4c42
feat(scripts): H3 prior-art runner: batch invariance, accuracy and la…
fusheng-ji Oct 8, 2026
64e4e4e
Merge the H3 prior-art runner into PR #484
fusheng-ji Oct 8, 2026
fbdee41
Merge the H3 prior-art runner into PR #485
fusheng-ji Oct 8, 2026
71af326
Merge the H3 prior-art runner into PR #486
fusheng-ji Oct 8, 2026
526f3ae
Merge the H3 prior-art runner into PR #487
fusheng-ji Oct 8, 2026
b179a80
Merge the H3 prior-art runner into PR #488
fusheng-ji Oct 8, 2026
0eb5b2d
Merge the H3 prior-art runner into PR #489
fusheng-ji Oct 8, 2026
d8ede13
Merge the H3 prior-art runner into PR #493
fusheng-ji Oct 8, 2026
464cbdc
Merge the H3 prior-art runner into PR #494
fusheng-ji Oct 8, 2026
98b6637
fix(scripts): keep h3_prior_art mode results next to the report, not …
fusheng-ji Oct 8, 2026
78b3c1c
Merge the H3 prior-art runner fix into PR #484
fusheng-ji Oct 8, 2026
04f8725
Merge the H3 prior-art runner fix into PR #485
fusheng-ji Oct 8, 2026
3900e6a
Merge the H3 prior-art runner fix into PR #486
fusheng-ji Oct 8, 2026
55faa06
Merge the H3 prior-art runner fix into PR #487
fusheng-ji Oct 8, 2026
d61153e
Merge the H3 prior-art runner fix into PR #488
fusheng-ji Oct 8, 2026
4ffb00f
Merge the H3 prior-art runner fix into PR #489
fusheng-ji Oct 8, 2026
b1bca0c
Merge the H3 prior-art runner fix into PR #493
fusheng-ji Oct 8, 2026
10e73eb
Merge the H3 prior-art runner fix into PR #494
fusheng-ji Oct 8, 2026
ae267aa
fix(scripts): name the failing check in the H3 prior-art figure
fusheng-ji Oct 8, 2026
c53c7bc
Merge the H3 prior-art figure fix into PR #484
fusheng-ji Oct 8, 2026
e09326c
Merge the H3 prior-art figure fix into PR #485
fusheng-ji Oct 8, 2026
a8999fd
Merge the H3 prior-art figure fix into PR #486
fusheng-ji Oct 8, 2026
3c291d8
Merge the H3 prior-art figure fix into PR #487
fusheng-ji Oct 8, 2026
fd8f223
Merge the H3 prior-art figure fix into PR #488
fusheng-ji Oct 8, 2026
189363b
Merge the H3 prior-art figure fix into PR #489
fusheng-ji Oct 8, 2026
846edfb
Merge the H3 prior-art figure fix into PR #493
fusheng-ji Oct 8, 2026
c1bb8ff
Merge the H3 prior-art figure fix into PR #494
fusheng-ji Oct 8, 2026
2a0a472
fix(scripts): shorter batch-invariance panel title in the H3 prior-ar…
fusheng-ji Oct 8, 2026
9baded5
Merge the H3 prior-art figure fix into PR #484
fusheng-ji Oct 8, 2026
a7e52f0
Merge the H3 prior-art figure fix into PR #485
fusheng-ji Oct 8, 2026
2cc97c3
Merge the H3 prior-art figure fix into PR #486
fusheng-ji Oct 8, 2026
a81c431
Merge the H3 prior-art figure fix into PR #487
fusheng-ji Oct 8, 2026
ce8ff11
Merge the H3 prior-art figure fix into PR #488
fusheng-ji Oct 8, 2026
06ce261
Merge the H3 prior-art figure fix into PR #489
fusheng-ji Oct 8, 2026
bbf7c8a
Merge the H3 prior-art figure fix into PR #493
fusheng-ji Oct 8, 2026
9536002
Merge the H3 prior-art figure fix into PR #494
fusheng-ji Oct 8, 2026
fd6cba6
docs(h3): adaln_row_gather vs existing implementations (B200, 6c900ae)
fusheng-ji Oct 8, 2026
05aa61f
docs(h3): gate_residual vs existing implementations (B200, 4010854)
fusheng-ji Oct 8, 2026
058b79b
docs(h3): timestep_sinusoid vs existing implementations (B200, 7a82917)
fusheng-ji Oct 8, 2026
6c04b08
docs(h3): timestep_mlp vs existing implementations (B200, ffd958a)
fusheng-ji Oct 8, 2026
b452e1a
docs(h3): adaln_projection vs existing implementations (B200, c88691c)
fusheng-ji Oct 8, 2026
51cac27
docs(h3): norm_modulate vs existing implementations (B200, ee83dec)
fusheng-ji Oct 8, 2026
8fe6ed0
docs(h3): final_adaln_out vs existing implementations (B200, 03da729)
fusheng-ji Oct 8, 2026
f58d337
Merge the H3 prior-art evidence into PR #484
fusheng-ji Oct 8, 2026
9726960
Merge the H3 prior-art evidence into PR #485
fusheng-ji Oct 8, 2026
fac1de5
Merge the H3 prior-art evidence into PR #486
fusheng-ji Oct 8, 2026
657ade7
Merge the H3 prior-art evidence into PR #487
fusheng-ji Oct 8, 2026
10c2d96
Merge the H3 prior-art evidence into PR #488
fusheng-ji Oct 8, 2026
df578b4
Merge the H3 prior-art evidence into PR #489
fusheng-ji Oct 8, 2026
9e6b686
Merge the H3 prior-art evidence into PR #493
fusheng-ji Oct 8, 2026
cbacb52
Merge the H3 prior-art evidence into PR #494
fusheng-ji Oct 8, 2026
5be1c31
fix(scripts): readable H3 prior-art figure when errors are close or l…
fusheng-ji Oct 8, 2026
5a4dd60
docs(h3): regenerate the timestep_sinusoid figure with the fixed layout
fusheng-ji Oct 8, 2026
6d08483
Merge the H3 prior-art figure layout fix into PR #484
fusheng-ji Oct 8, 2026
4543e36
docs(h3): regenerate the timestep_mlp figure with the fixed layout
fusheng-ji Oct 8, 2026
8e2a687
Merge the H3 prior-art figure layout fix into PR #485
fusheng-ji Oct 8, 2026
3d15509
docs(h3): regenerate the adaln_projection figure with the fixed layout
fusheng-ji Oct 8, 2026
c870650
Merge the H3 prior-art figure layout fix into PR #486
fusheng-ji Oct 8, 2026
a8f3b3d
docs(h3): regenerate the adaln_row_gather figure with the fixed layout
fusheng-ji Oct 8, 2026
f1b6df3
Merge the H3 prior-art figure layout fix into PR #487
fusheng-ji Oct 8, 2026
43e1fd3
docs(h3): regenerate the norm_modulate figure with the fixed layout
fusheng-ji Oct 8, 2026
c1829ec
Merge the H3 prior-art figure layout fix into PR #488
fusheng-ji Oct 8, 2026
8123f57
docs(h3): regenerate the gate_residual figure with the fixed layout
fusheng-ji Oct 8, 2026
6d87d01
Merge the H3 prior-art figure layout fix into PR #489
fusheng-ji Oct 8, 2026
9056d7e
docs(h3): regenerate the final_adaln_out figure with the fixed layout
fusheng-ji Oct 8, 2026
ac60da2
Merge the H3 prior-art figure layout fix into PR #493
fusheng-ji Oct 8, 2026
074e34c
Merge the H3 prior-art figure layout fix into PR #494
fusheng-ji Oct 8, 2026
6a4d42a
Fix H3 native validation, backward timing, and rank cleanup
fusheng-ji Oct 9, 2026
fbee9a5
Regenerate gate and final AdaLN backward-only B200 evidence
fusheng-ji Oct 9, 2026
101edfa
test(h3): consolidate native validation tests for review file limit
fusheng-ji Oct 9, 2026
26d4992
fix(h3): address automated review bounds, device, and benchmark findings
fusheng-ji Oct 9, 2026
fc0a3d5
Merge test-h3 refactor and migrate H3 timestep validation
fusheng-ji Oct 10, 2026
74a20fb
Merge test-h3 refactor and migrate H3 timestep MLP validation
fusheng-ji Oct 10, 2026
0c94934
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
1fe5fdf
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
f13d1b1
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
45bc42b
Merge migrated H3 dependency PR #483
fusheng-ji Oct 10, 2026
35dccee
Merge migrated H3 dependency PR #484
fusheng-ji Oct 10, 2026
479e778
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
562efb5
Merge migrated H3 dependency PR #485
fusheng-ji Oct 10, 2026
7faa68e
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
cc0b590
Merge test-h3 refactor and migrate H3 SP kernels, tools, and tests
fusheng-ji Oct 10, 2026
38c6e65
Merge migrated H3 dependency PR #486
fusheng-ji Oct 10, 2026
6b6ea54
Merge migrated H3 dependency PR #487
fusheng-ji Oct 10, 2026
9885264
Merge migrated H3 dependency PR #488
fusheng-ji Oct 10, 2026
484923f
Merge test-h3 refactor and migrate H3 kernels, tools, and tests
fusheng-ji Oct 10, 2026
4f5112c
Merge migrated H3 dependency PR #489
fusheng-ji Oct 10, 2026
3f191f1
Merge migrated H3 TP dependency and retain SP review fixes
fusheng-ji Oct 10, 2026
8f860d2
Sync migrated H3 stack ancestry
fusheng-ji Oct 10, 2026
97b42de
docs(h3): point row-gather validation to migrated tests
fusheng-ji Oct 10, 2026
01a37f7
Merge migrated validation documentation from PR #486
fusheng-ji Oct 10, 2026
293b232
Merge migrated validation documentation from PR #487
fusheng-ji Oct 10, 2026
1bfe8f5
Merge migrated validation documentation from PR #488
fusheng-ji Oct 10, 2026
9c4c1fc
Merge migrated validation documentation from PR #489
fusheng-ji Oct 10, 2026
dfdd984
Merge migrated validation documentation from PR #493
fusheng-ji Oct 10, 2026
41e4293
feat(h3): one real MiniMax-H3 block, forward and backward, node by node
fusheng-ji Oct 11, 2026
caaf754
docs(h3): one-block evidence on B200 (41e4293) and operator doc
fusheng-ji Oct 11, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions .coderabbit.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
# yaml-language-server: $schema=https://coderabbit.ai/integrations/schema.v2.json
inheritance: true
reviews:
auto_review:
base_branches:
- "^test-h3$"
71 changes: 71 additions & 0 deletions benchmarks/models/benchmark_h3_conditioning.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright (c) 2026 RL-Kernel Contributors

"""Benchmark the MiniMax-H3 conditioning ops (RFC #420) against the provider path.

CUDA-event medians, bandwidth and peak memory per case, plus the backend
the registry dispatched. Timings alternate candidate/provider execution order;
backward timings exclude forward setup. Cases live in ``rl_engine/validation/models/h3_report.py``.

python benchmarks/models/benchmark_h3_conditioning.py --op timestep_sinusoid_h3
python benchmarks/models/benchmark_h3_conditioning.py --op all --json out.json
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[2]))

import torch # noqa: E402

from rl_engine.runtime.registry import KernelRegistry # noqa: E402
from rl_engine.validation.models.h3_chain import environment # noqa: E402
from rl_engine.validation.models.h3_report import PERF_CASES, TIMED_KEYS, measure # noqa: E402


def main() -> None:
"""Benchmark selected CUDA operators, print summaries, and optionally save JSON."""

parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--op", default="all", choices=["all", *PERF_CASES])
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--iters", type=int, default=200)
parser.add_argument("--json", type=Path, default=None)
args = parser.parse_args()
if not torch.cuda.is_available():
raise SystemExit("needs a CUDA device")
torch.backends.cuda.matmul.allow_tf32 = False

registry = KernelRegistry()
env = environment()
print(f"device={env['gpu']} torch={env['torch']} cuda={env['cuda']}")
print("timing order alternates each iteration; backward timings exclude forward setup")
results = []
for name in list(PERF_CASES) if args.op == "all" else [args.op]:
for case in PERF_CASES[name](registry):
row = measure(case, args.warmup, args.iters)
results.append(row)
parts = [
f"{key}={row[f'{key}_us']:.2f}us ({row[f'{key}_gbps']:.1f} GB/s, "
f"peak {row[f'{key}_peak_mib']:.2f} MiB)"
for key in TIMED_KEYS
if f"{key}_us" in row
]
order = row["execution_order"]
phases = [" -> ".join(order[f"iteration_{i}"]) for i in (0, 1)]
print(
f"{name} {row['case']} [{row['backend']}]: "
+ "; ".join(parts)
+ "; alternating order: "
+ " / ".join(phases)
)
if args.json is not None:
args.json.write_text(json.dumps({"environment": env, "results": results}, indent=2) + "\n")


if __name__ == "__main__":
main()
5 changes: 5 additions & 0 deletions build_tools/extensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,11 @@ def get_extensions():
# CUDA IPC and the fixed-tree collective implementation are not
# part of the ROCm extension.
cuda_sources.append("csrc/cuda/collectives/deterministic_collective.cu")
cuda_sources.append("csrc/cuda/h3/timestep_sinusoid.cu")
cuda_sources.append("csrc/cuda/h3/det_linear.cu")
cuda_sources.append("csrc/cuda/h3/adaln_row_gather.cu")
cuda_sources.append("csrc/cuda/h3/rmsnorm_modulate.cu")
cuda_sources.append("csrc/cuda/h3/gate_residual.cu")
# This source contains NVIDIA PTX (cp.async, ldmatrix, and mma.sync).
# The ROCm dispatcher falls back to PyTorch SDPA for this operator.
cuda_sources.append("csrc/cuda/attention/prefix_shared_attention.cu")
Expand Down
134 changes: 134 additions & 0 deletions csrc/bindings/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -500,6 +500,67 @@ at::Tensor prefix_shared_attention(
#endif
#endif

// MiniMax-H3 (RFC #420) conditioning-path declarations. CUDA only.
#if !defined(USE_ROCM) && !defined(KERNEL_ALIGN_WITH_ROCM) && \
(defined(__CUDACC__) || defined(KERNEL_ALIGN_WITH_CUDA))
torch::Tensor h3_timestep_sinusoid_forward(torch::Tensor timestep, int64_t num_channels,
double max_period, bool check_range);
std::vector<torch::Tensor> h3_det_linear_forward(torch::Tensor x, torch::Tensor weight,
c10::optional<torch::Tensor> bias,
int64_t activation, bool save_pre_activation);
torch::Tensor h3_det_linear_backward_input(torch::Tensor grad, torch::Tensor weight,
c10::ScalarType out_dtype);
torch::Tensor h3_det_linear_backward_input_partials(torch::Tensor grad, torch::Tensor weight);
torch::Tensor h3_det_linear_fold_chunks(torch::Tensor partial, c10::ScalarType out_dtype);
std::vector<torch::Tensor> h3_det_linear_backward_weight(torch::Tensor grad, torch::Tensor x,
c10::ScalarType w_dtype,
bool with_bias);
torch::Tensor h3_adaln_row_gather_forward(torch::Tensor rows, torch::Tensor timestep_indices,
torch::Tensor token_tags, int64_t chunks,
int64_t modality_num);
torch::Tensor h3_adaln_row_gather_backward(torch::Tensor grad, torch::Tensor sorted_pos,
torch::Tensor tile_begin, torch::Tensor tile_end,
torch::Tensor seg_first_tile,
c10::ScalarType out_dtype);
std::vector<torch::Tensor> h3_rmsnorm_forward(torch::Tensor x, torch::Tensor weight, double eps,
c10::optional<torch::Tensor> shift,
c10::optional<torch::Tensor> scale,
c10::optional<torch::Tensor> index);
std::vector<torch::Tensor> h3_rmsnorm_backward(
torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd,
c10::optional<torch::Tensor> shift, c10::optional<torch::Tensor> scale,
c10::optional<torch::Tensor> index, c10::optional<torch::Tensor> sorted_pos,
c10::optional<torch::Tensor> tile_begin, c10::optional<torch::Tensor> tile_end,
c10::optional<torch::Tensor> seg_first_tile);
torch::Tensor h3_gate_residual_forward(torch::Tensor residual, torch::Tensor y, torch::Tensor gate,
torch::Tensor index);
std::vector<torch::Tensor> h3_gate_residual_backward(torch::Tensor grad, torch::Tensor y,
torch::Tensor gate, torch::Tensor index,
torch::Tensor sorted_pos,
torch::Tensor tile_begin,
torch::Tensor tile_end,
torch::Tensor seg_first_tile);
std::vector<torch::Tensor> h3_rmsnorm_backward_partials(
torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd,
c10::optional<torch::Tensor> shift, c10::optional<torch::Tensor> scale,
c10::optional<torch::Tensor> index, torch::Tensor dw_rows, torch::Tensor dw_begin,
torch::Tensor dw_end, c10::optional<torch::Tensor> seg_rows,
c10::optional<torch::Tensor> seg_begin, c10::optional<torch::Tensor> seg_end);
std::vector<torch::Tensor> h3_rmsnorm_fold_partials(torch::Tensor dw_partial, torch::Tensor weight,
c10::optional<torch::Tensor> seg_partial,
c10::optional<torch::Tensor> seg_first_tile);
torch::Tensor h3_gate_grad_partials(torch::Tensor grad, torch::Tensor y, torch::Tensor rows,
torch::Tensor tile_begin, torch::Tensor tile_end);
torch::Tensor h3_gate_grad_fold(torch::Tensor partial, torch::Tensor seg_first_tile,
c10::ScalarType dtype);
torch::Tensor h3_rmsnorm_backward_dx(torch::Tensor grad, torch::Tensor x, torch::Tensor weight,
torch::Tensor rstd, c10::optional<torch::Tensor> shift,
c10::optional<torch::Tensor> scale,
c10::optional<torch::Tensor> index);
torch::Tensor h3_gate_residual_backward_dy(torch::Tensor grad, torch::Tensor gate,
torch::Tensor index);
#endif

// PyBind11 Module Registration
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "RL-Kernel High-Performance Operator Extension Library";
Expand Down Expand Up @@ -758,4 +819,77 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
"Deterministic GPT-NeoX token-major RoPE apply for ROCm");
#endif
#endif

#if !defined(USE_ROCM) && !defined(KERNEL_ALIGN_WITH_ROCM) && \
(defined(__CUDACC__) || defined(KERNEL_ALIGN_WITH_CUDA))
m.def("h3_timestep_sinusoid_forward", torch::wrap_pybind_function(h3_timestep_sinusoid_forward),
"MiniMax-H3 FP32 [cos | sin] timestep features, bitwise to the diffusers CUDA path",
py::arg("timestep"), py::arg("num_channels") = 256, py::arg("max_period") = 10000.0,
py::arg("check_range") = true);
m.def("h3_det_linear_forward", &h3_det_linear_forward,
"Batch-invariant warp-per-column linear (contract h3-det-linear-v1), optional SiLU",
py::arg("x"), py::arg("weight"), py::arg("bias") = py::none(),
py::arg("activation") = 0, py::arg("save_pre_activation") = false);
m.def("h3_det_linear_backward_input", &h3_det_linear_backward_input,
"Deterministic grad @ weight with fixed 64-row N chunks folded in order",
py::arg("grad"), py::arg("weight"), py::arg("out_dtype"));
m.def("h3_det_linear_backward_input_partials", &h3_det_linear_backward_input_partials,
"Per-64-row-chunk FP32 partials of grad @ weight, before the ascending fold",
py::arg("grad"), py::arg("weight"));
m.def("h3_det_linear_fold_chunks", &h3_det_linear_fold_chunks,
"Ascending left fold of h3_det_linear_backward_input_partials, cast once",
py::arg("partial"), py::arg("out_dtype"));
m.def("h3_det_linear_backward_weight", &h3_det_linear_backward_weight,
"Deterministic dW/dbias as ascending-row FP32 folds",
py::arg("grad"), py::arg("x"), py::arg("w_dtype"), py::arg("with_bias") = true);
m.def("h3_adaln_row_gather_forward", &h3_adaln_row_gather_forward,
"Fused six-way AdaLN row gather by timestep_index * 3 + token_tag (pure copy)",
py::arg("rows"), py::arg("timestep_indices"), py::arg("token_tags"),
py::arg("chunks") = 6, py::arg("modality_num") = 3);
m.def("h3_adaln_row_gather_backward", &h3_adaln_row_gather_backward,
"Deterministic segmented sum (sorted tiles folded in order) for the row gather",
py::arg("grad"), py::arg("sorted_pos"), py::arg("tile_begin"), py::arg("tile_end"),
py::arg("seg_first_tile"), py::arg("out_dtype"));
m.def("h3_rmsnorm_forward", &h3_rmsnorm_forward,
"RMSNorm replaying PyTorch's reduction order, with optional fused AdaLN modulation",
py::arg("x"), py::arg("weight"), py::arg("eps"), py::arg("shift") = py::none(),
py::arg("scale") = py::none(), py::arg("index") = py::none());
m.def("h3_rmsnorm_backward", &h3_rmsnorm_backward,
"Deterministic RMSNorm(+modulation) backward: row-local dx, tiled dweight, sorted table grads",
py::arg("grad"), py::arg("x"), py::arg("weight"), py::arg("rstd"),
py::arg("shift") = py::none(), py::arg("scale") = py::none(),
py::arg("index") = py::none(), py::arg("sorted_pos") = py::none(),
py::arg("tile_begin") = py::none(), py::arg("tile_end") = py::none(),
py::arg("seg_first_tile") = py::none());
m.def("h3_gate_residual_forward", &h3_gate_residual_forward,
"residual + gate[index] * y with the gate row gathered in-kernel (eager rounding order)",
py::arg("residual"), py::arg("y"), py::arg("gate"), py::arg("index"));
m.def("h3_gate_residual_backward", &h3_gate_residual_backward,
"Gated-residual backward: exact dy, deterministic sorted segment sum for dgate",
py::arg("grad"), py::arg("y"), py::arg("gate"), py::arg("index"), py::arg("sorted_pos"),
py::arg("tile_begin"), py::arg("tile_end"), py::arg("seg_first_tile"));
m.def("h3_rmsnorm_backward_partials", &h3_rmsnorm_backward_partials,
"WS1 RMSNorm dweight / table-gradient tile partials over explicit row lists (SP)",
py::arg("grad"), py::arg("x"), py::arg("weight"), py::arg("rstd"), py::arg("shift"),
py::arg("scale"), py::arg("index"), py::arg("dw_rows"), py::arg("dw_begin"),
py::arg("dw_end"), py::arg("seg_rows") = py::none(), py::arg("seg_begin") = py::none(),
py::arg("seg_end") = py::none());
m.def("h3_rmsnorm_fold_partials", &h3_rmsnorm_fold_partials,
"WS1 ascending folds of RMSNorm dweight and per-segment table-gradient partials",
py::arg("dw_partial"), py::arg("weight"), py::arg("seg_partial") = py::none(),
py::arg("seg_first_tile") = py::none());
m.def("h3_gate_grad_partials", &h3_gate_grad_partials,
"WS1 d_gate tile partials over explicit row lists (SP)", py::arg("grad"), py::arg("y"),
py::arg("rows"), py::arg("tile_begin"), py::arg("tile_end"));
m.def("h3_gate_grad_fold", &h3_gate_grad_fold,
"WS1 per-segment ascending fold of d_gate partials, cast once", py::arg("partial"),
py::arg("seg_first_tile"), py::arg("dtype"));
m.def("h3_rmsnorm_backward_dx", &h3_rmsnorm_backward_dx,
"Row-local dx of the WS1 RMSNorm(+modulation) backward", py::arg("grad"), py::arg("x"),
py::arg("weight"), py::arg("rstd"), py::arg("shift") = py::none(),
py::arg("scale") = py::none(), py::arg("index") = py::none());
m.def("h3_gate_residual_backward_dy", &h3_gate_residual_backward_dy,
"Row-local d_y = grad * gate[index] of the WS1 gated-residual backward",
py::arg("grad"), py::arg("gate"), py::arg("index"));
#endif
}
Loading