diff --git a/.coderabbit.yaml b/.coderabbit.yaml new file mode 100644 index 000000000..ebea58ae2 --- /dev/null +++ b/.coderabbit.yaml @@ -0,0 +1,6 @@ +# yaml-language-server: $schema=https://coderabbit.ai/integrations/schema.v2.json +inheritance: true +reviews: + auto_review: + base_branches: + - "^test-h3$" diff --git a/benchmarks/models/benchmark_h3_conditioning.py b/benchmarks/models/benchmark_h3_conditioning.py new file mode 100644 index 000000000..261556809 --- /dev/null +++ b/benchmarks/models/benchmark_h3_conditioning.py @@ -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() diff --git a/build_tools/extensions.py b/build_tools/extensions.py index db079e828..b3bade305 100644 --- a/build_tools/extensions.py +++ b/build_tools/extensions.py @@ -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") diff --git a/csrc/bindings/ops.cpp b/csrc/bindings/ops.cpp index 0c7c3e7f4..9322e23d0 100644 --- a/csrc/bindings/ops.cpp +++ b/csrc/bindings/ops.cpp @@ -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 h3_det_linear_forward(torch::Tensor x, torch::Tensor weight, + c10::optional 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 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 h3_rmsnorm_forward(torch::Tensor x, torch::Tensor weight, double eps, + c10::optional shift, + c10::optional scale, + c10::optional index); +std::vector h3_rmsnorm_backward( + torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, + c10::optional shift, c10::optional scale, + c10::optional index, c10::optional sorted_pos, + c10::optional tile_begin, c10::optional tile_end, + c10::optional seg_first_tile); +torch::Tensor h3_gate_residual_forward(torch::Tensor residual, torch::Tensor y, torch::Tensor gate, + torch::Tensor index); +std::vector 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 h3_rmsnorm_backward_partials( + torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, + c10::optional shift, c10::optional scale, + c10::optional index, torch::Tensor dw_rows, torch::Tensor dw_begin, + torch::Tensor dw_end, c10::optional seg_rows, + c10::optional seg_begin, c10::optional seg_end); +std::vector h3_rmsnorm_fold_partials(torch::Tensor dw_partial, torch::Tensor weight, + c10::optional seg_partial, + c10::optional 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 shift, + c10::optional scale, + c10::optional 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"; @@ -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 } diff --git a/csrc/cuda/h3/adaln_row_gather.cu b/csrc/cuda/h3/adaln_row_gather.cu new file mode 100644 index 000000000..0de2d98d5 --- /dev/null +++ b/csrc/cuda/h3/adaln_row_gather.cu @@ -0,0 +1,259 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors +// +// MiniMax-H3 AdaLN row gather (RFC #420 `adaln_row_gather`). +// +// rows: (3T, C * H) modulation rows, row r = timestep * 3 + modality and +// column block c = shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, +// gate_mlp. For every packed-sequence position s: +// +// r[s] = timestep_indices[s] * modality_num + token_tags[s] +// out[c, s, :] = rows[r[s], c * H : (c + 1) * H] +// +// Forward is a pure copy: bitwise equal to C index_select calls, one launch, +// independent of S, packing order and repeats. +// +// Backward d_rows[r, j] = sum_{s : r[s] = r} grad[c(j), s, h(j)] is the one +// cross-row reduction. Positions are visited in a fixed order: the caller +// sorts them by (r, s) with a stable sort and cuts every segment into tiles of +// at most kTile positions. Each tile is an ascending FP32 chain from 0; the +// tiles of a segment are then left-folded in ascending order and cast once. +// No atomics: the result is a function of (r[s], grad) only. + +#include +#include +#include +#include +#include + +#include + +namespace { + +__device__ __forceinline__ float to_float(float v) { return v; } +__device__ __forceinline__ float to_float(__nv_bfloat16 v) { return __bfloat162float(v); } +__device__ __forceinline__ float to_float(__half v) { return __half2float(v); } +__device__ __forceinline__ float to_float(double v) { return static_cast(v); } + +template +__device__ __forceinline__ T from_float(float v); +template <> +__device__ __forceinline__ float from_float(float v) { return v; } +template <> +__device__ __forceinline__ __nv_bfloat16 from_float<__nv_bfloat16>(float v) { + return __float2bfloat16(v); +} +template <> +__device__ __forceinline__ __half from_float<__half>(float v) { return __float2half(v); } +template <> +__device__ __forceinline__ double from_float(float v) { return v; } + +template +struct CudaType { + using type = T; +}; +template <> +struct CudaType { + using type = __nv_bfloat16; +}; +template <> +struct CudaType { + using type = __half; +}; + +template +__device__ __forceinline__ int64_t adaln_row(const idx_t* ti, const idx_t* tags, int64_t s, + int64_t modality_num, int64_t num_timesteps) { + const int64_t timestep = static_cast(ti[s]); + const int64_t tag = static_cast(tags[s]); + // Check each semantic index before multiplication, including values that overflow int64. + CUDA_KERNEL_ASSERT(timestep >= 0 && timestep < num_timesteps); + CUDA_KERNEL_ASSERT(tag >= 0 && tag < modality_num); + return timestep * modality_num + tag; +} + +// One block per (s, c); 16-byte copies when every row start is 16-byte aligned. +template +__global__ void adaln_row_gather_kernel(const T* __restrict__ rows, int64_t row_stride, + const idx_t* __restrict__ ti, + const idx_t* __restrict__ tags, T* __restrict__ out, + int64_t seq, int64_t hidden, int64_t chunks, + int64_t modality_num, int64_t num_timesteps) { + for (int64_t s = blockIdx.x; s < seq; s += gridDim.x) { + const int64_t c = blockIdx.y; + const int64_t r = adaln_row(ti, tags, s, modality_num, num_timesteps); + const T* src = rows + r * row_stride + c * hidden; + T* dst = out + (c * seq + s) * hidden; + if constexpr (kVector) { + constexpr int kVec = 16 / sizeof(T); + const uint4* src4 = reinterpret_cast(src); + uint4* dst4 = reinterpret_cast(dst); + for (int64_t i = threadIdx.x; i < hidden / kVec; i += blockDim.x) dst4[i] = src4[i]; + } else { + for (int64_t i = threadIdx.x; i < hidden; i += blockDim.x) dst[i] = src[i]; + } + } +} + +// partial[tile, j] = sum over sorted positions p in tile (ascending) of grad[c, s_p, h] +template +__global__ void adaln_row_gather_partial_kernel(const g_t* __restrict__ grad, + const int64_t* __restrict__ sorted_pos, + const int64_t* __restrict__ tile_begin, + const int64_t* __restrict__ tile_end, + float* __restrict__ partial, int64_t seq, + int64_t hidden, int64_t width) { + const int64_t tile = blockIdx.x; + const int64_t j = static_cast(blockIdx.y) * blockDim.x + threadIdx.x; + if (j >= width) return; + const int64_t c = j / hidden; + const int64_t h = j - c * hidden; + const g_t* g = grad + c * seq * hidden + h; + float acc = 0.0f; + const int64_t begin = tile_begin[tile], end = tile_end[tile]; + CUDA_KERNEL_ASSERT(begin >= 0 && begin <= end && end <= seq); + for (int64_t p = begin; p < end; ++p) { + CUDA_KERNEL_ASSERT(sorted_pos[p] >= 0 && sorted_pos[p] < seq); + acc += to_float(g[sorted_pos[p] * hidden]); + } + partial[tile * width + j] = acc; +} + +template +__global__ void adaln_row_gather_fold_kernel(const float* __restrict__ partial, + const int64_t* __restrict__ seg_first_tile, + out_t* __restrict__ out, int64_t width, int64_t tiles) { + const int64_t r = blockIdx.x; + const int64_t j = static_cast(blockIdx.y) * blockDim.x + threadIdx.x; + if (j >= width) return; + float acc = 0.0f; + const int64_t begin = seg_first_tile[r], end = seg_first_tile[r + 1]; + CUDA_KERNEL_ASSERT(begin >= 0 && begin <= end && end <= tiles); + for (int64_t tile = begin; tile < end; ++tile) { + acc += partial[tile * width + j]; + } + out[r * width + j] = from_float(acc); +} + +void check_index(const torch::Tensor& t, const char* name, int64_t seq, + const torch::Tensor& like) { + TORCH_CHECK(t.is_cuda() && t.device() == like.device(), name, " must be on ", like.device()); + TORCH_CHECK(t.dim() == 1 && t.size(0) == seq, name, " must be (", seq, ",)"); + TORCH_CHECK(t.scalar_type() == at::kLong || t.scalar_type() == at::kInt, name, + " must be int64 or int32"); + TORCH_CHECK(t.is_contiguous(), name, " must be contiguous"); +} + +} // namespace + +// rows (R, C * H) with unit column stride -> out (C, S, H) contiguous. +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_CHECK(rows.is_cuda() && rows.dim() == 2 && rows.stride(1) == 1, + "rows must be a 2-D CUDA tensor with unit column stride"); + TORCH_CHECK(chunks > 0 && rows.size(1) % chunks == 0, "rows width ", rows.size(1), + " is not a multiple of ", chunks); + TORCH_CHECK(modality_num > 0 && rows.size(0) % modality_num == 0, "rows count ", + rows.size(0), " is not a multiple of ", modality_num); + const int64_t num_timesteps = rows.size(0) / modality_num; + const int64_t seq = timestep_indices.numel(); + TORCH_CHECK(seq > 0, "the packed sequence must not be empty"); + check_index(timestep_indices, "timestep_indices", seq, rows); + check_index(token_tags, "token_tags", seq, rows); + TORCH_CHECK(timestep_indices.scalar_type() == token_tags.scalar_type(), + "timestep_indices and token_tags must share an integer dtype"); + const int64_t hidden = rows.size(1) / chunks; + const c10::cuda::CUDAGuard device_guard(rows.device()); + auto out = torch::empty({chunks, seq, hidden}, rows.options()); + auto stream = at::cuda::getCurrentCUDAStream(); + const dim3 grid(static_cast(std::min(seq, 1 << 20)), + static_cast(chunks)); + const int threads = 128; + const int64_t elem = rows.element_size(); + const bool vector = (hidden * elem) % 16 == 0 && (rows.stride(0) * elem) % 16 == 0 && + reinterpret_cast(rows.data_ptr()) % 16 == 0; + + AT_DISPATCH_FLOATING_TYPES_AND2( + at::kHalf, at::kBFloat16, rows.scalar_type(), "h3_adaln_row_gather_forward", [&] { + using T = typename CudaType::type; + const T* src = reinterpret_cast(rows.data_ptr()); + T* dst = reinterpret_cast(out.data_ptr()); + auto launch = [&](auto idx_tag) { + using idx_t = decltype(idx_tag); + const idx_t* ti = timestep_indices.data_ptr(); + const idx_t* tags = token_tags.data_ptr(); + if (vector) { + adaln_row_gather_kernel<<>>( + src, rows.stride(0), ti, tags, dst, seq, hidden, chunks, modality_num, + num_timesteps); + } else { + adaln_row_gather_kernel<<>>( + src, rows.stride(0), ti, tags, dst, seq, hidden, chunks, modality_num, + num_timesteps); + } + }; + if (timestep_indices.scalar_type() == at::kLong) { + launch(int64_t{0}); + } else { + launch(int32_t{0}); + } + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} + +// grad (C, S, H) contiguous; sorted_pos (S,) positions sorted stably by row; +// tiles [tile_begin, tile_end) in sorted order; seg_first_tile (R + 1,). +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) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 3 && grad.is_contiguous(), + "grad must be a contiguous (C, S, H) CUDA tensor"); + for (const auto* t : {&sorted_pos, &tile_begin, &tile_end, &seg_first_tile}) { + TORCH_CHECK(t->is_cuda() && t->device() == grad.device() && + t->dim() == 1 && t->is_contiguous() && + t->scalar_type() == at::kLong, + "tile metadata must be contiguous 1-D int64 tensors on grad device"); + } + TORCH_CHECK(tile_begin.numel() == tile_end.numel(), "tile_begin/tile_end length mismatch"); + const int64_t chunks = grad.size(0); + const int64_t seq = grad.size(1); + const int64_t hidden = grad.size(2); + const int64_t width = chunks * hidden; + const int64_t num_rows = seg_first_tile.numel() - 1; + const int64_t tiles = tile_begin.numel(); + TORCH_CHECK(chunks > 0 && seq > 0 && hidden > 0 && num_rows > 0, + "gradient dimensions and table row count must be positive"); + TORCH_CHECK(sorted_pos.numel() == seq, "sorted_pos must have S entries"); + const c10::cuda::CUDAGuard device_guard(grad.device()); + auto out = torch::empty({num_rows, width}, grad.options().dtype(out_dtype)); + auto partial = torch::empty({std::max(tiles, 1), width}, + grad.options().dtype(at::kFloat)); + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + const unsigned col_blocks = static_cast((width + threads - 1) / threads); + if (tiles > 0) { + AT_DISPATCH_FLOATING_TYPES_AND2( + at::kHalf, at::kBFloat16, grad.scalar_type(), "h3_adaln_row_gather_partial", [&] { + using G = typename CudaType::type; + adaln_row_gather_partial_kernel + <<(tiles), col_blocks), threads, 0, stream>>>( + reinterpret_cast(grad.data_ptr()), + sorted_pos.data_ptr(), tile_begin.data_ptr(), + tile_end.data_ptr(), partial.data_ptr(), seq, hidden, width); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + AT_DISPATCH_FLOATING_TYPES_AND2( + at::kHalf, at::kBFloat16, out_dtype, "h3_adaln_row_gather_fold", [&] { + using O = typename CudaType::type; + adaln_row_gather_fold_kernel + <<(num_rows), col_blocks), threads, 0, stream>>>( + partial.data_ptr(), seg_first_tile.data_ptr(), + reinterpret_cast(out.data_ptr()), width, tiles); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} diff --git a/csrc/cuda/h3/det_linear.cu b/csrc/cuda/h3/det_linear.cu new file mode 100644 index 000000000..96b954b0e --- /dev/null +++ b/csrc/cuda/h3/det_linear.cu @@ -0,0 +1,620 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors +// +// Deterministic, batch-invariant row-wise linear layers for the MiniMax-H3 +// conditioning path (RFC #420: `timestep_mlp_fp32`, `adaln_projection_3mod`). +// +// These layers run on a handful of rows (one per distinct timestep) against +// large weights, so they are GEMVs bound by weight bandwidth. The reduction +// order is fixed and depends only on K (and N for the input gradient), never +// on the number of rows or their position. +// +// Reduction contract h3-det-linear-v1 +// ----------------------------------- +// forward y[t, n] = act(bias[n] + dot(x[t, :], W[n, :])) +// * one warp per output column n; lane l owns the 16-byte K chunks +// c = l, l + 32, l + 64, ... (VEC = 16 / sizeof(w) elements each); +// * each lane accumulates its chunks in ascending order and the elements of +// a chunk in ascending order with fmaf into an FP32 accumulator from 0; +// * the 32 lane sums are combined with an xor butterfly (16, 8, 4, 2, 1), +// whose result is identical on every lane (FP add is commutative); +// * the bias is added once after the butterfly; the activation runs in FP32 +// (SiLU as v / (1 + expf(-v)), PyTorch's formula); one cast at the store. +// * this FP32 forward is used for FP32 weights; BF16 weights take the +// tensor-core forward of contract h3-det-linear-bf16-mma-v1 below. +// d_input dx[t, k] = sum_n g[t, n] * W[n, k] +// * N is split into fixed chunks of 64 rows; each chunk is an ascending +// fmaf chain into FP32 from 0; the chunk partials are then left-folded in +// ascending chunk order and cast once. +// The partials and the fold are also bound separately: a contiguous N +// slice starting on a chunk boundary yields exactly the full call's +// partials for those chunks, which is what tensor-parallel shards gather. +// d_weight dW[n, k] = sum_t g[t, n] * x[t, k]; d_bias db[n] = sum_t g[t, n] +// * an ascending-t fmaf chain into FP32 from 0, cast once. This is the one +// cross-row reduction; its order is the logical row order. +// +// Reduction contract h3-det-linear-bf16-mma-v1 (BF16 x and W, FP32 accumulate) +// ----------------------------------------------------------------------------- +// forward y[t, n] = act(bias[n] + dot(x[t, :], W[n, :])) +// * a warp owns 16 output columns (MMA rows) and 8 input rows (MMA columns; +// rows past T are zero), so every launch runs the same instruction +// sequence for every T <= 8, and an MMA column never sees another's data; +// * K is visited in groups of 16 in ascending order. Within a 32-wide step, +// quad q supplies k = 32s + 8q + {0..3} to the first mma.sync m16n8k16 and +// k = 32s + 8q + {4..7} to the second (the same k map for W and x); +// * every mma.sync starts from a zero accumulator, so the tensor core only +// sums 16 products; its FP32 result is added to the running FP32 sum with +// an IEEE add, group by group in ascending k; +// * then bias, FP32 activation and one cast, as above. + +#include +#include +#include +#include +#include + +#include +#include + +namespace { + +#ifndef H3_FWD_F32_COLS +#define H3_FWD_F32_COLS 2 +#define H3_FWD_F32_AHEAD 6 +#endif +#ifndef H3_FWD_CONFIGS +#define H3_FWD_CONFIGS H3_FWD_CASE(H3_FWD_F32_COLS, H3_FWD_F32_AHEAD) +#endif + +constexpr int kWarp = 32; +constexpr int kWarpsPerBlock = 8; +constexpr int kRowTile = 4; +constexpr int kDInputChunk = 64; + +enum Activation : int64_t { kActNone = 0, kActSilu = 1 }; + +__device__ __forceinline__ float to_float(float v) { return v; } +__device__ __forceinline__ float to_float(__nv_bfloat16 v) { return __bfloat162float(v); } + +template +__device__ __forceinline__ T from_float(float v); +template <> +__device__ __forceinline__ float from_float(float v) { return v; } +template <> +__device__ __forceinline__ __nv_bfloat16 from_float<__nv_bfloat16>(float v) { + return __float2bfloat16(v); // round to nearest even, like at::BFloat16 +} + +// 16-byte vector load converted to FP32. +template +struct Vec16 { + static constexpr int kN = 16 / sizeof(T); + __device__ __forceinline__ static void load(const T* ptr, float (&out)[kN]) { + unpack(load_raw(ptr), out); + } + __device__ __forceinline__ static uint4 load_raw(const T* ptr) { + return *reinterpret_cast(ptr); + } + // Streaming weights are read once: bypass L1 allocation. + __device__ __forceinline__ static uint4 load_stream(const T* ptr) { + return __ldcs(reinterpret_cast(ptr)); + } + __device__ __forceinline__ static void unpack(const uint4& raw, float (&out)[kN]) { + const T* vals = reinterpret_cast(&raw); +#pragma unroll + for (int i = 0; i < kN; ++i) out[i] = to_float(vals[i]); + } +}; + +__device__ __forceinline__ float warp_butterfly_sum(float v) { +#pragma unroll + for (int offset = kWarp / 2; offset > 0; offset >>= 1) { + v += __shfl_xor_sync(0xffffffffu, v, offset); + } + return v; +} + +// Each warp owns kCols adjacent output columns so that kLoadAhead x kCols +// 16-byte weight loads are in flight per lane. Grouping columns and chunks +// changes only when loads are issued: every column still accumulates its own +// chunks in ascending order, exactly as a one-column-per-warp kernel would. +template +__global__ void __launch_bounds__(kWarpsPerBlock * kWarp) + det_linear_forward_kernel(const x_t* __restrict__ x, const w_t* __restrict__ w, + const w_t* __restrict__ bias, out_t* __restrict__ out, + float* __restrict__ pre_act, int64_t rows, int64_t n_out, + int64_t k_in, int64_t activation) { + using WV = Vec16; + using XV = Vec16; + constexpr int kVec = WV::kN; + static_assert(kCols * kRows <= kWarp, "one warp cannot store more than kWarp outputs per tile"); + static_assert(XV::kN == kVec, "x and weight share a dtype"); + constexpr int64_t kStride = static_cast(kWarp) * kVec; + + const int lane = threadIdx.x % kWarp; + const int64_t warp_global = + static_cast(blockIdx.x) * kWarpsPerBlock + threadIdx.x / kWarp; + const int64_t warp_stride = static_cast(gridDim.x) * kWarpsPerBlock; + + for (int64_t n0 = warp_global * kCols; n0 < n_out; n0 += warp_stride * kCols) { + for (int64_t t0 = 0; t0 < rows; t0 += kRows) { + float acc[kCols][kRows]; +#pragma unroll + for (int c = 0; c < kCols; ++c) { +#pragma unroll + for (int r = 0; r < kRows; ++r) acc[c][r] = 0.0f; + } + for (int64_t g0 = static_cast(lane) * kVec; g0 < k_in; + g0 += kStride * kLoadAhead) { + uint4 raw[kLoadAhead][kCols]; +#pragma unroll + for (int u = 0; u < kLoadAhead; ++u) { + const int64_t k0 = g0 + u * kStride; +#pragma unroll + for (int c = 0; c < kCols; ++c) { + if (k0 < k_in && n0 + c < n_out) raw[u][c] = WV::load_stream(w + (n0 + c) * k_in + k0); + } + } +#pragma unroll + for (int u = 0; u < kLoadAhead; ++u) { + const int64_t k0 = g0 + u * kStride; + if (k0 >= k_in) break; +#pragma unroll + for (int r = 0; r < kRows; ++r) { + if (t0 + r < rows) { + float xv[kVec]; + XV::load(x + (t0 + r) * k_in + k0, xv); +#pragma unroll + for (int c = 0; c < kCols; ++c) { + float wv[kVec]; + WV::unpack(raw[u][c], wv); +#pragma unroll + for (int j = 0; j < kVec; ++j) acc[c][r] = fmaf(xv[j], wv[j], acc[c][r]); + } + } + } + } + } +#pragma unroll + for (int c = 0; c < kCols; ++c) { + const int64_t n = n0 + c; +#pragma unroll + for (int r = 0; r < kRows; ++r) { + const int64_t t = t0 + r; + const float sum = warp_butterfly_sum(acc[c][r]); + if (lane == c * kRows + r && t < rows && n < n_out) { + float v = sum + (bias != nullptr ? to_float(bias[n]) : 0.0f); + if (pre_act != nullptr) pre_act[t * n_out + n] = v; + if (activation == kActSilu) v = v / (1.0f + expf(-v)); + out[t * n_out + n] = from_float(v); + } + } + } + } + } +} + +__device__ __forceinline__ void mma_m16n8k16_bf16(float (&c)[4], const uint32_t (&a)[4], + uint32_t b0, uint32_t b1) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ < 800 + // BF16 mma.sync needs SM80+; the host refuses to launch below that. + __trap(); +#else + asm volatile( + "mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0,%1,%2,%3}, {%4,%5,%6,%7}, " + "{%8,%9}, {%0,%1,%2,%3};\n" + : "+f"(c[0]), "+f"(c[1]), "+f"(c[2]), "+f"(c[3]) + : "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b0), "r"(b1)); +#endif +} + +constexpr int kMmaCols = 16; // output columns per warp (MMA M) +constexpr int kMmaRows = 8; // input rows per pass (MMA N) +constexpr int kMmaAhead = 4; // 32-wide K steps loaded before use + +__device__ __forceinline__ void mma_group_accumulate(float (&acc)[4], const uint32_t (&a)[4], + uint32_t b0, uint32_t b1) { + float part[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + mma_m16n8k16_bf16(part, a, b0, b1); +#pragma unroll + for (int i = 0; i < 4; ++i) acc[i] += part[i]; +} + +// Contract h3-det-linear-bf16-mma-v1. Lane = 4 * g + q: g picks the MMA row +// (output columns n0 + g and n0 + g + 8) and the MMA column (input row t0 + g). +__global__ void __launch_bounds__(kWarpsPerBlock * kWarp) + det_linear_forward_bf16_mma_kernel(const __nv_bfloat16* __restrict__ x, + const __nv_bfloat16* __restrict__ w, + const __nv_bfloat16* __restrict__ bias, + __nv_bfloat16* __restrict__ out, float* __restrict__ pre_act, + int64_t rows, int64_t n_out, int64_t k_in, + int64_t activation) { + const int lane = threadIdx.x % kWarp; + const int g = lane >> 2; + const int q = lane & 3; + const int64_t warp_global = + static_cast(blockIdx.x) * kWarpsPerBlock + threadIdx.x / kWarp; + const int64_t warp_stride = static_cast(gridDim.x) * kWarpsPerBlock; + const uint4 zero = make_uint4(0u, 0u, 0u, 0u); + + for (int64_t n0 = warp_global * kMmaCols; n0 < n_out; n0 += warp_stride * kMmaCols) { + const bool lo_ok = n0 + g < n_out; + const bool hi_ok = n0 + g + 8 < n_out; + const __nv_bfloat16* w_lo = w + (n0 + g) * k_in; + const __nv_bfloat16* w_hi = w + (n0 + g + 8) * k_in; + for (int64_t t0 = 0; t0 < rows; t0 += kMmaRows) { + const bool x_ok = t0 + g < rows; + const __nv_bfloat16* x_row = x + (t0 + g) * k_in; + float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f}; + for (int64_t kb = 0; kb < k_in; kb += 32 * kMmaAhead) { + uint4 a_lo[kMmaAhead], a_hi[kMmaAhead], bx[kMmaAhead]; +#pragma unroll + for (int u = 0; u < kMmaAhead; ++u) { + const int64_t k = kb + 32 * u + 8 * q; + const bool k_ok = k < k_in; + a_lo[u] = (k_ok && lo_ok) ? __ldcs(reinterpret_cast(w_lo + k)) : zero; + a_hi[u] = (k_ok && hi_ok) ? __ldcs(reinterpret_cast(w_hi + k)) : zero; + bx[u] = (k_ok && x_ok) ? *reinterpret_cast(x_row + k) : zero; + } +#pragma unroll + for (int u = 0; u < kMmaAhead; ++u) { + if (kb + 32 * u >= k_in) break; + const uint32_t first[4] = {a_lo[u].x, a_hi[u].x, a_lo[u].y, a_hi[u].y}; + const uint32_t second[4] = {a_lo[u].z, a_hi[u].z, a_lo[u].w, a_hi[u].w}; + mma_group_accumulate(acc, first, bx[u].x, bx[u].y); + mma_group_accumulate(acc, second, bx[u].z, bx[u].w); + } + } + // acc = {(n0+g, t0+2q), (n0+g, t0+2q+1), (n0+g+8, t0+2q), (n0+g+8, t0+2q+1)} +#pragma unroll + for (int i = 0; i < 4; ++i) { + const int64_t n = n0 + g + (i >= 2 ? 8 : 0); + const int64_t t = t0 + 2 * q + (i & 1); + if (n < n_out && t < rows) { + float v = acc[i] + (bias != nullptr ? __bfloat162float(bias[n]) : 0.0f); + if (pre_act != nullptr) pre_act[t * n_out + n] = v; + if (activation == kActSilu) v = v / (1.0f + expf(-v)); + out[t * n_out + n] = __float2bfloat16(v); + } + } + } + } +} + +// partial[c, t, k] = sum_{n in chunk c} g[t, n] * w[n, k] +template +__global__ void det_linear_dinput_partial_kernel(const g_t* __restrict__ grad, + const w_t* __restrict__ w, + float* __restrict__ partial, int64_t rows, + int64_t n_out, int64_t k_in) { + const int64_t k = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t chunk = blockIdx.y; + if (k >= k_in) return; + const int64_t n_begin = chunk * kDInputChunk; + const int64_t n_end = min(n_begin + kDInputChunk, n_out); + for (int64_t t0 = 0; t0 < rows; t0 += kRowTile) { + float acc[kRowTile]; +#pragma unroll + for (int r = 0; r < kRowTile; ++r) acc[r] = 0.0f; + for (int64_t n = n_begin; n < n_end; ++n) { + const float wv = to_float(w[n * k_in + k]); +#pragma unroll + for (int r = 0; r < kRowTile; ++r) { + const int64_t t = t0 + r; + if (t < rows) acc[r] = fmaf(to_float(grad[t * n_out + n]), wv, acc[r]); + } + } +#pragma unroll + for (int r = 0; r < kRowTile; ++r) { + const int64_t t = t0 + r; + if (t < rows) partial[(chunk * rows + t) * k_in + k] = acc[r]; + } + } +} + +template +__global__ void det_fold_chunks_kernel(const float* __restrict__ partial, out_t* __restrict__ out, + int64_t chunks, int64_t elems) { + const int64_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (i >= elems) return; + float acc = partial[i]; + for (int64_t c = 1; c < chunks; ++c) acc += partial[c * elems + i]; + out[i] = from_float(acc); +} + +// dW[n, k] = sum_t g[t, n] * x[t, k] ; ascending t, FP32 from 0, one cast. +template +__global__ void det_linear_dweight_kernel(const g_t* __restrict__ grad, const x_t* __restrict__ x, + w_t* __restrict__ dw, int64_t rows, int64_t n_out, + int64_t k_in) { + const int64_t total = n_out * k_in; + for (int64_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; i < total; + i += static_cast(gridDim.x) * blockDim.x) { + const int64_t n = i / k_in; + const int64_t k = i - n * k_in; + float acc = 0.0f; + for (int64_t t = 0; t < rows; ++t) { + acc = fmaf(to_float(grad[t * n_out + n]), to_float(x[t * k_in + k]), acc); + } + dw[i] = from_float(acc); + } +} + +template +__global__ void det_linear_dbias_kernel(const g_t* __restrict__ grad, w_t* __restrict__ db, + int64_t rows, int64_t n_out) { + const int64_t n = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (n >= n_out) return; + float acc = 0.0f; + for (int64_t t = 0; t < rows; ++t) acc += to_float(grad[t * n_out + n]); + db[n] = from_float(acc); +} + +template +struct CudaType { + using type = T; +}; +template <> +struct CudaType { + using type = __nv_bfloat16; +}; + +template +const typename CudaType::type* cptr(const torch::Tensor& t) { + return reinterpret_cast::type*>(t.data_ptr()); +} +template +typename CudaType::type* mptr(torch::Tensor& t) { + return reinterpret_cast::type*>(t.data_ptr()); +} + +void check_matrix(const torch::Tensor& t, const char* name) { + TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor"); + TORCH_CHECK(t.dim() == 2, name, " must be 2-D"); + TORCH_CHECK(t.is_contiguous(), name, " must be contiguous"); + TORCH_CHECK(t.scalar_type() == at::kFloat || t.scalar_type() == at::kBFloat16, name, + " must be float32 or bfloat16, got ", t.scalar_type()); + TORCH_CHECK(reinterpret_cast(t.data_ptr()) % 16 == 0, name, + " must be 16-byte aligned"); +} + +int64_t grid_for(int64_t work, int64_t per_block) { + const int64_t blocks = (work + per_block - 1) / per_block; + return std::max(1, std::min(blocks, 1 << 20)); +} + +// Launch shape only: the row tile, the columns per warp and the load depth +// change which loads are in flight, never a column's accumulation order. +struct ForwardConfig { + int rows; + int cols; + int ahead; +}; + +ForwardConfig pick_forward_config(int64_t rows) { + const int row_tile = rows <= 1 ? 1 : (rows <= 2 ? 2 : 4); + return {row_tile, H3_FWD_F32_COLS, H3_FWD_F32_AHEAD}; +} + +template +void launch_forward_impl(const T* x, const T* w, const T* bias, T* out, float* pre, + int64_t rows, int64_t n_out, int64_t k_in, int64_t activation, + cudaStream_t stream) { + const int64_t blocks = grid_for((n_out + kCols - 1) / kCols, kWarpsPerBlock); + det_linear_forward_kernel + <<(blocks), kWarpsPerBlock * kWarp, 0, stream>>>( + x, w, bias, out, pre, rows, n_out, k_in, activation); +} + +template +void launch_forward_rows(const ForwardConfig& cfg, const T* x, const T* w, const T* bias, T* out, + float* pre, int64_t rows, int64_t n_out, int64_t k_in, + int64_t activation, cudaStream_t stream) { +#define H3_FWD_CASE(C, A) \ + if (cfg.cols == C && cfg.ahead == A) { \ + launch_forward_impl(x, w, bias, out, pre, rows, n_out, k_in, activation, \ + stream); \ + return; \ + } + H3_FWD_CONFIGS +#undef H3_FWD_CASE + TORCH_CHECK(false, "unsupported det_linear forward config cols=", cfg.cols, + " ahead=", cfg.ahead); +} + +template +void launch_forward(const ForwardConfig& cfg, const T* x, const T* w, const T* bias, T* out, + float* pre, int64_t rows, int64_t n_out, int64_t k_in, int64_t activation, + cudaStream_t stream) { + if (cfg.rows == 1) { + launch_forward_rows(cfg, x, w, bias, out, pre, rows, n_out, k_in, activation, stream); + } else if (cfg.rows == 2) { + launch_forward_rows(cfg, x, w, bias, out, pre, rows, n_out, k_in, activation, stream); + } else { + launch_forward_rows(cfg, x, w, bias, out, pre, rows, n_out, k_in, activation, stream); + } +} + +} // namespace + +std::vector h3_det_linear_forward(torch::Tensor x, torch::Tensor weight, + c10::optional bias, + int64_t activation, bool save_pre_activation) { + check_matrix(x, "x"); + check_matrix(weight, "weight"); + TORCH_CHECK(x.device() == weight.device(), "x and weight must be on the same device"); + TORCH_CHECK(x.scalar_type() == weight.scalar_type(), + "x and weight must share a dtype (cast at the declared boundary first), got ", + x.scalar_type(), " and ", weight.scalar_type()); + TORCH_CHECK(activation == kActNone || activation == kActSilu, "unknown activation ", + activation); + const int64_t rows = x.size(0); + const int64_t k_in = x.size(1); + const int64_t n_out = weight.size(0); + TORCH_CHECK(rows > 0, "x must have at least one row"); + TORCH_CHECK(weight.size(1) == k_in, "weight is [", n_out, ", ", weight.size(1), + "] but x has K=", k_in); + const int64_t vec = 16 / x.element_size(); + TORCH_CHECK(k_in % vec == 0, "K=", k_in, " must be a multiple of ", vec, + " for 16-byte vector loads"); + if (bias.has_value()) { + TORCH_CHECK(bias->is_cuda() && bias->dim() == 1 && bias->size(0) == n_out && + bias->is_contiguous() && bias->scalar_type() == weight.scalar_type(), + "bias must be a contiguous [N] CUDA tensor with the weight dtype"); + TORCH_CHECK(bias->device() == weight.device(), + "bias and weight must be on the same device"); + } + + const c10::cuda::CUDAGuard device_guard(x.device()); + auto out = torch::empty({rows, n_out}, x.options()); + torch::Tensor pre; + if (save_pre_activation) pre = torch::empty({rows, n_out}, x.options().dtype(at::kFloat)); + auto stream = at::cuda::getCurrentCUDAStream(); + if (x.scalar_type() == at::kFloat) { + const ForwardConfig cfg = pick_forward_config(rows); + launch_forward(cfg, cptr(x), cptr(weight), + bias.has_value() ? cptr(*bias) : nullptr, mptr(out), + save_pre_activation ? pre.data_ptr() : nullptr, rows, n_out, + k_in, activation, stream); + } else { + TORCH_CHECK(at::cuda::getCurrentDeviceProperties()->major >= 8, + "the BF16 h3 det_linear forward uses mma.sync and needs SM80 or newer"); + const int64_t blocks = grid_for((n_out + kMmaCols - 1) / kMmaCols, kWarpsPerBlock); + det_linear_forward_bf16_mma_kernel<<(blocks), kWarpsPerBlock * kWarp, + 0, stream>>>( + cptr(x), cptr(weight), + bias.has_value() ? cptr(*bias) : nullptr, mptr(out), + save_pre_activation ? pre.data_ptr() : nullptr, rows, n_out, k_in, activation); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + if (save_pre_activation) return {out, pre}; + return {out}; +} + +// grad [T, N] (float32), weight [N, K] -> partial [ceil(N / 64), T, K] float32: +// the per-chunk sums of contract h3-det-linear-v1, before the ascending fold. +// A contiguous slice of N that starts on a chunk boundary yields exactly the +// matching rows of the full call (tensor-parallel shards rely on this). +torch::Tensor h3_det_linear_backward_input_partials(torch::Tensor grad, torch::Tensor weight) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 2 && grad.is_contiguous() && + grad.scalar_type() == at::kFloat, + "grad must be a contiguous 2-D float32 CUDA tensor"); + check_matrix(weight, "weight"); + TORCH_CHECK(grad.device() == weight.device(), "grad and weight must be on the same device"); + TORCH_CHECK(grad.size(1) == weight.size(0), "grad N != weight N"); + const int64_t rows = grad.size(0); + TORCH_CHECK(rows > 0, "grad must have at least one row"); + const int64_t n_out = weight.size(0); + const int64_t k_in = weight.size(1); + const c10::cuda::CUDAGuard device_guard(grad.device()); + const int64_t chunks = (n_out + kDInputChunk - 1) / kDInputChunk; + auto partial = torch::empty({chunks, rows, k_in}, grad.options()); + if (partial.numel() == 0) return partial; + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + dim3 grid(static_cast((k_in + threads - 1) / threads), static_cast(chunks)); + if (weight.scalar_type() == at::kFloat) { + det_linear_dinput_partial_kernel<<>>( + grad.data_ptr(), cptr(weight), partial.data_ptr(), rows, n_out, + k_in); + } else { + det_linear_dinput_partial_kernel<<>>( + grad.data_ptr(), cptr(weight), partial.data_ptr(), rows, + n_out, k_in); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return partial; +} + +// partial [C, T, K] float32 -> left fold over C in ascending order, cast once. +torch::Tensor h3_det_linear_fold_chunks(torch::Tensor partial, c10::ScalarType out_dtype) { + TORCH_CHECK(partial.is_cuda() && partial.dim() == 3 && partial.is_contiguous() && + partial.scalar_type() == at::kFloat && partial.size(0) > 0, + "partial must be a non-empty contiguous 3-D float32 CUDA tensor"); + TORCH_CHECK(out_dtype == at::kFloat || out_dtype == at::kBFloat16, + "out_dtype must be float32 or bfloat16"); + const c10::cuda::CUDAGuard device_guard(partial.device()); + const int64_t chunks = partial.size(0); + auto out = torch::empty({partial.size(1), partial.size(2)}, partial.options().dtype(out_dtype)); + const int64_t elems = out.numel(); + if (elems == 0) return out; + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + const unsigned fold_blocks = static_cast((elems + threads - 1) / threads); + if (out_dtype == at::kFloat) { + det_fold_chunks_kernel<<>>( + partial.data_ptr(), mptr(out), chunks, elems); + } else { + det_fold_chunks_kernel<__nv_bfloat16><<>>( + partial.data_ptr(), mptr(out), chunks, elems); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} + +// grad [T, N] (float32), weight [N, K] -> grad_input [T, K] in out_dtype. +torch::Tensor h3_det_linear_backward_input(torch::Tensor grad, torch::Tensor weight, + c10::ScalarType out_dtype) { + TORCH_CHECK(out_dtype == at::kFloat || out_dtype == at::kBFloat16, + "out_dtype must be float32 or bfloat16"); + auto partial = h3_det_linear_backward_input_partials(grad, weight); + if (partial.size(0) == 0) { // N == 0: no output column contributes + return torch::zeros({partial.size(1), partial.size(2)}, grad.options().dtype(out_dtype)); + } + return h3_det_linear_fold_chunks(partial, out_dtype); +} + +// grad [T, N] float32, x [T, K] (float32 or bfloat16) -> (dW [N, K], dbias [N]) in w_dtype. +std::vector h3_det_linear_backward_weight(torch::Tensor grad, torch::Tensor x, + c10::ScalarType w_dtype, + bool with_bias) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 2 && grad.is_contiguous() && + grad.scalar_type() == at::kFloat, + "grad must be a contiguous 2-D float32 CUDA tensor"); + check_matrix(x, "x"); + TORCH_CHECK(grad.device() == x.device(), "grad and x must be on the same device"); + TORCH_CHECK(grad.size(0) == x.size(0), "grad rows != x rows"); + TORCH_CHECK(w_dtype == at::kFloat || w_dtype == at::kBFloat16, + "w_dtype must be float32 or bfloat16"); + const int64_t rows = grad.size(0); + TORCH_CHECK(rows > 0, "grad must have at least one row"); + const int64_t n_out = grad.size(1); + const int64_t k_in = x.size(1); + const c10::cuda::CUDAGuard device_guard(grad.device()); + auto dw = torch::empty({n_out, k_in}, grad.options().dtype(w_dtype)); + if (n_out == 0) { + if (!with_bias) return {dw}; + return {dw, torch::empty({n_out}, grad.options().dtype(w_dtype))}; + } + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + const int64_t total = n_out * k_in; + const unsigned blocks = static_cast(grid_for(total, threads)); + const float* g = grad.data_ptr(); +#define H3_DW_LAUNCH(XT, XC, WT, WC) \ + det_linear_dweight_kernel<<>>( \ + g, cptr(x), mptr(dw), rows, n_out, k_in) + if (x.scalar_type() == at::kFloat && w_dtype == at::kFloat) { + H3_DW_LAUNCH(float, float, float, float); + } else if (x.scalar_type() == at::kFloat) { + H3_DW_LAUNCH(float, float, at::BFloat16, __nv_bfloat16); + } else if (w_dtype == at::kFloat) { + H3_DW_LAUNCH(at::BFloat16, __nv_bfloat16, float, float); + } else { + H3_DW_LAUNCH(at::BFloat16, __nv_bfloat16, at::BFloat16, __nv_bfloat16); + } +#undef H3_DW_LAUNCH + C10_CUDA_KERNEL_LAUNCH_CHECK(); + if (!with_bias) return {dw}; + auto db = torch::empty({n_out}, grad.options().dtype(w_dtype)); + const unsigned bias_blocks = static_cast((n_out + threads - 1) / threads); + if (w_dtype == at::kFloat) { + det_linear_dbias_kernel<<>>( + g, mptr(db), rows, n_out); + } else { + det_linear_dbias_kernel<<>>( + g, mptr(db), rows, n_out); + } + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {dw, db}; +} diff --git a/csrc/cuda/h3/gate_residual.cu b/csrc/cuda/h3/gate_residual.cu new file mode 100644 index 000000000..eb0b03f50 --- /dev/null +++ b/csrc/cuda/h3/gate_residual.cu @@ -0,0 +1,326 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors +// +// MiniMax-H3 gated residual (RFC #420 `adaln_gate_residual`). +// +// out[r, j] = cast(residual[r, j] + cast(gate[i, j] * y[r, j])) i = index[r % S] +// +// The gate row is read from a strided view of the AdaLN table (gate_msa or +// gate_mlp) inside the kernel, and the two roundings sit exactly where the +// eager `residual + gate.index_select(0, i) * y` rounds (__fmul_rn/__fadd_rn +// keep FP32 from contracting them into one FMA), so the result is +// bitwise equal to diffusers. Elementwise: every output depends only on its +// own inputs. +// +// Backward: d_residual = grad; d_y = cast(grad * gate) (the product of two +// 16-bit values is exact in FP32, so this equals the eager VJP bitwise); +// d_gate[i] = sum over positions mapped to i of grad * y, an FP32 segmented +// sum over positions sorted stably by row, in fixed tiles folded in order. +// The tile partials and the fold are also bound separately for sequence +// parallel callers (each tile computed where its rows live, folded in order). + +#include +#include +#include +#include +#include +#include + +#include +#include + +namespace { + +__device__ __forceinline__ float to_f(float v) { return v; } +__device__ __forceinline__ float to_f(__nv_bfloat16 v) { return __bfloat162float(v); } +__device__ __forceinline__ float to_f(__half v) { return __half2float(v); } + +template +__device__ __forceinline__ T from_f(float v); +template <> +__device__ __forceinline__ float from_f(float v) { return v; } +template <> +__device__ __forceinline__ __nv_bfloat16 from_f<__nv_bfloat16>(float v) { return __float2bfloat16(v); } +template <> +__device__ __forceinline__ __half from_f<__half>(float v) { return __float2half(v); } + +template +__device__ __forceinline__ float round_t(float v) { return to_f(from_f(v)); } + +struct Gate { + const void* rows; // (R, N) view, row stride `stride` + int64_t stride; + const int64_t* index; // (S,) + int64_t seq; + int64_t num_rows; +}; + +// One block per row (grid-strided); 16-byte vectors when the row allows it. +// `kVecElems` elements are loaded together; each is computed independently. +template +__global__ void gate_residual_rows_kernel(const T* __restrict__ a, const T* __restrict__ b, + T* __restrict__ out, int64_t rows, int64_t n, + Gate gate, bool vectorized) { + constexpr int kVecElems = 16 / sizeof(T); + for (int64_t r = blockIdx.x; r < rows; r += gridDim.x) { + const int64_t i = gate.index[r % gate.seq]; + CUDA_KERNEL_ASSERT(i >= 0 && i < gate.num_rows); + const T* g = static_cast(gate.rows) + i * gate.stride; + const T* ar = a + r * n; + const T* br = kDy ? nullptr : b + r * n; + T* orow = out + r * n; + if (vectorized) { + for (int64_t v = threadIdx.x; v < n / kVecElems; v += blockDim.x) { + const uint4 av = reinterpret_cast(ar)[v]; + const uint4 gv = reinterpret_cast(g)[v]; + uint4 bv = av; + if (!kDy) bv = reinterpret_cast(br)[v]; + const T* ae = reinterpret_cast(&av); + const T* ge = reinterpret_cast(&gv); + const T* be = reinterpret_cast(&bv); + uint4 ov; + T* oe = reinterpret_cast(&ov); +#pragma unroll + for (int k = 0; k < kVecElems; ++k) { + if (kDy) { + oe[k] = from_f(to_f(ae[k]) * to_f(ge[k])); // grad * gate + } else { + const float p = round_t(__fmul_rn(to_f(ge[k]), to_f(be[k]))); // gate * y + oe[k] = from_f(__fadd_rn(to_f(ae[k]), p)); // residual + p + } + } + reinterpret_cast(orow)[v] = ov; + } + } else { + for (int64_t j = threadIdx.x; j < n; j += blockDim.x) { + if (kDy) { + orow[j] = from_f(to_f(ar[j]) * to_f(g[j])); + } else { + const float p = round_t(__fmul_rn(to_f(g[j]), to_f(br[j]))); + orow[j] = from_f(__fadd_rn(to_f(ar[j]), p)); + } + } + } + } +} + +// partial[tile, j] = sum over sorted positions of the tile (ascending) of grad * y +template +__global__ void gate_grad_partial_kernel(const T* __restrict__ grad, const T* __restrict__ y, + const int64_t* __restrict__ sorted_pos, + const int64_t* __restrict__ tile_begin, + const int64_t* __restrict__ tile_end, + float* __restrict__ partial, int64_t n) { + const int64_t tile = blockIdx.x; + const int64_t j = static_cast(blockIdx.y) * blockDim.x + threadIdx.x; + if (j >= n) return; + float acc = 0.0f; + for (int64_t p = tile_begin[tile]; p < tile_end[tile]; ++p) { + const int64_t e = sorted_pos[p] * n + j; + acc = fmaf(to_f(grad[e]), to_f(y[e]), acc); + } + partial[tile * n + j] = acc; +} + +template +__global__ void gate_grad_fold_kernel(const float* __restrict__ partial, + const int64_t* __restrict__ seg_first_tile, + T* __restrict__ out, int64_t n) { + const int64_t seg = blockIdx.x; + const int64_t j = static_cast(blockIdx.y) * blockDim.x + threadIdx.x; + if (j >= n) return; + float acc = 0.0f; + for (int64_t t = seg_first_tile[seg]; t < seg_first_tile[seg + 1]; ++t) acc += partial[t * n + j]; + out[seg * n + j] = from_f(acc); +} + +#define H3_DISPATCH(TYPE, NAME, ...) \ + [&] { \ + switch (TYPE) { \ + case at::kFloat: { using T = float; __VA_ARGS__(); break; } \ + case at::kHalf: { using T = __half; __VA_ARGS__(); break; } \ + case at::kBFloat16: { using T = __nv_bfloat16; __VA_ARGS__(); break; } \ + default: TORCH_CHECK(false, NAME, ": unsupported dtype ", TYPE); \ + } \ + }() + +Gate make_gate(const torch::Tensor& gate, const torch::Tensor& index, const torch::Tensor& x) { + TORCH_CHECK(gate.is_cuda() && gate.device() == x.device() && gate.scalar_type() == x.scalar_type(), + "gate must be a CUDA tensor with the activations' dtype"); + TORCH_CHECK(gate.dim() == 2 && gate.size(1) == x.size(1) && gate.stride(1) == 1, + "gate must be (R, N) with unit column stride"); + TORCH_CHECK(index.is_cuda() && index.device() == x.device() && index.scalar_type() == at::kLong && + index.dim() == 1 && index.is_contiguous() && index.numel() > 0, + "index must be a non-empty contiguous int64 tensor on ", x.device()); + TORCH_CHECK(x.size(0) % index.size(0) == 0, "rows must be a multiple of S"); + return Gate{gate.data_ptr(), gate.stride(0), index.data_ptr(), index.size(0), gate.size(0)}; +} + +void check_act(const torch::Tensor& t, const char* name, const torch::Tensor& like) { + TORCH_CHECK(t.is_cuda() && t.device() == like.device() && t.scalar_type() == like.scalar_type(), + name, " must match the residual's device and dtype"); + TORCH_CHECK(t.sizes() == like.sizes() && t.is_contiguous(), name, + " must be contiguous with the residual's shape"); +} + +unsigned row_blocks(int64_t rows) { + return static_cast(std::min(rows, 1 << 30)); +} + +bool can_vectorize(const torch::Tensor& gate, int64_t n, std::initializer_list acts) { + const int64_t vec = 16 / gate.element_size(); + if (n % vec != 0 || (gate.stride(0) * gate.element_size()) % 16 != 0) return false; + if (reinterpret_cast(gate.data_ptr()) % 16 != 0) return false; + for (const auto* t : acts) { + if (reinterpret_cast(t->data_ptr()) % 16 != 0) return false; + } + return true; +} + +} // namespace + +// residual, y: (M, N); gate: (R, N) view; index: (S,), M % S == 0. +torch::Tensor h3_gate_residual_forward(torch::Tensor residual, torch::Tensor y, torch::Tensor gate, + torch::Tensor index) { + TORCH_CHECK(residual.is_cuda() && residual.dim() == 2 && residual.is_contiguous() && + residual.numel() > 0, + "residual must be a non-empty contiguous (M, N) CUDA tensor"); + check_act(y, "y", residual); + const Gate g = make_gate(gate, index, residual); + const c10::cuda::CUDAGuard guard(residual.device()); + auto out = torch::empty_like(residual); + auto stream = at::cuda::getCurrentCUDAStream(); + const bool vec = can_vectorize(gate, residual.size(1), {&residual, &y, &out}); + H3_DISPATCH(residual.scalar_type(), "h3_gate_residual_forward", [&] { + gate_residual_rows_kernel<<>>( + reinterpret_cast(residual.data_ptr()), reinterpret_cast(y.data_ptr()), + reinterpret_cast(out.data_ptr()), residual.size(0), residual.size(1), g, vec); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} + +// Returns {d_y, d_gate (R, N) in the gate dtype}. d_residual is the incoming grad. +std::vector 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) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 2 && grad.is_contiguous(), + "grad must be a contiguous (M, N) CUDA tensor"); + check_act(y, "y", grad); + const Gate g = make_gate(gate, index, grad); + for (const auto* t : {&sorted_pos, &tile_begin, &tile_end, &seg_first_tile}) { + TORCH_CHECK(t->is_cuda() && t->device() == grad.device() && t->scalar_type() == at::kLong && + t->dim() == 1 && t->is_contiguous(), + "tile metadata must be contiguous int64 tensors on ", grad.device()); + } + const int64_t rows = grad.size(0); + const int64_t n = grad.size(1); + const int64_t segments = seg_first_tile.numel() - 1; + TORCH_CHECK(segments == gate.size(0), "segments must equal the gate rows"); + const int64_t tiles = tile_begin.numel(); + const c10::cuda::CUDAGuard guard(grad.device()); + auto stream = at::cuda::getCurrentCUDAStream(); + auto dy = torch::empty_like(y); + auto partial = torch::empty({std::max(tiles, 1), n}, grad.options().dtype(at::kFloat)); + auto dgate = torch::empty({segments, n}, grad.options()); + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + H3_DISPATCH(grad.scalar_type(), "h3_gate_residual_backward", [&] { + const T* gp = reinterpret_cast(grad.data_ptr()); + const T* yp = reinterpret_cast(y.data_ptr()); + const bool vec = can_vectorize(gate, n, {&grad, &dy}); + gate_residual_rows_kernel<<>>( + gp, nullptr, reinterpret_cast(dy.data_ptr()), rows, n, g, vec); + if (tiles > 0) { + gate_grad_partial_kernel<<(tiles), col_blocks), threads, 0, stream>>>( + gp, yp, sorted_pos.data_ptr(), tile_begin.data_ptr(), + tile_end.data_ptr(), partial.data_ptr(), n); + } + gate_grad_fold_kernel<<(segments), col_blocks), threads, 0, stream>>>( + partial.data_ptr(), seg_first_tile.data_ptr(), + reinterpret_cast(dgate.data_ptr()), n); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {dy, dgate}; +} + +// Sequence parallel: the WS1 d_gate tile partials over explicit row lists of +// (grad, y) rows -> (tiles, N) FP32. +torch::Tensor h3_gate_grad_partials(torch::Tensor grad, torch::Tensor y, torch::Tensor rows, + torch::Tensor tile_begin, torch::Tensor tile_end) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 2 && grad.is_contiguous(), + "grad must be a contiguous (M, N) CUDA tensor"); + check_act(y, "y", grad); + for (const auto* t : {&rows, &tile_begin, &tile_end}) { + TORCH_CHECK(t->is_cuda() && t->device() == grad.device() && t->scalar_type() == at::kLong && + t->dim() == 1 && t->is_contiguous(), + "tile metadata must be contiguous int64 CUDA tensors on grad's device"); + } + TORCH_CHECK(tile_begin.numel() == tile_end.numel(), "tile_begin and tile_end differ in length"); + const int64_t n = grad.size(1); + const int64_t tiles = tile_begin.numel(); + const c10::cuda::CUDAGuard guard(grad.device()); + auto partial = torch::empty({tiles, n}, grad.options().dtype(at::kFloat)); + if (tiles == 0) return partial; + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + H3_DISPATCH(grad.scalar_type(), "h3_gate_grad_partials", [&] { + gate_grad_partial_kernel<<(tiles), col_blocks), threads, 0, + at::cuda::getCurrentCUDAStream()>>>( + reinterpret_cast(grad.data_ptr()), reinterpret_cast(y.data_ptr()), + rows.data_ptr(), tile_begin.data_ptr(), tile_end.data_ptr(), + partial.data_ptr(), n); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return partial; +} + +// The WS1 d_gate fold: each segment's tiles in ascending order, cast once to dtype. +torch::Tensor h3_gate_grad_fold(torch::Tensor partial, torch::Tensor seg_first_tile, + c10::ScalarType dtype) { + TORCH_CHECK(partial.is_cuda() && partial.dim() == 2 && partial.is_contiguous() && + partial.scalar_type() == at::kFloat, + "partial must be a contiguous float32 (tiles, N) CUDA tensor"); + TORCH_CHECK(seg_first_tile.is_cuda() && seg_first_tile.device() == partial.device() && + seg_first_tile.scalar_type() == at::kLong && seg_first_tile.dim() == 1 && + seg_first_tile.is_contiguous(), + "seg_first_tile must be a contiguous int64 CUDA tensor on partial's device"); + const int64_t n = partial.size(1); + const int64_t segments = seg_first_tile.numel() - 1; + const c10::cuda::CUDAGuard guard(partial.device()); + auto out = torch::empty({segments, n}, partial.options().dtype(dtype)); + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + H3_DISPATCH(dtype, "h3_gate_grad_fold", [&] { + gate_grad_fold_kernel<<(segments), col_blocks), threads, 0, + at::cuda::getCurrentCUDAStream()>>>( + partial.data_ptr(), seg_first_tile.data_ptr(), + reinterpret_cast(out.data_ptr()), n); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} + +// Row-local part of the WS1 backward only: d_y = cast(grad * gate[index]). +torch::Tensor h3_gate_residual_backward_dy(torch::Tensor grad, torch::Tensor gate, + torch::Tensor index) { + TORCH_CHECK(grad.is_cuda() && grad.dim() == 2 && grad.is_contiguous(), + "grad must be a contiguous (M, N) CUDA tensor"); + const Gate g = make_gate(gate, index, grad); + const int64_t rows = grad.size(0); + const int64_t n = grad.size(1); + const c10::cuda::CUDAGuard guard(grad.device()); + auto dy = torch::empty_like(grad); + H3_DISPATCH(grad.scalar_type(), "h3_gate_residual_backward_dy", [&] { + const bool vec = can_vectorize(gate, n, {&grad, &dy}); + gate_residual_rows_kernel<<>>( + reinterpret_cast(grad.data_ptr()), nullptr, reinterpret_cast(dy.data_ptr()), + rows, n, g, vec); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return dy; +} diff --git a/csrc/cuda/h3/rmsnorm_modulate.cu b/csrc/cuda/h3/rmsnorm_modulate.cu new file mode 100644 index 000000000..f9304ab13 --- /dev/null +++ b/csrc/cuda/h3/rmsnorm_modulate.cu @@ -0,0 +1,593 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors +// +// MiniMax-H3 RMSNorm and fused AdaLN modulation (RFC #420 `h3_rmsnorm`). +// +// n[r] = cast(w * (rstd[r] * x[r])) rstd = rsqrtf(sum(x^2) / N + eps) +// out[r] = cast(cast(n[r] * cast(1 + scale[i])) + shift[i]) i = row_index[r % S] +// +// Forward contract h3-rmsnorm-v1: the statistics replay PyTorch's +// vectorized_layer_norm_kernel (torch cf30153): one +// (32, 4) block per row, 4-element vectors, thread t sums vectors t, t+128, ... +// in order, a shuffle-down tree (16..1), a cross-warp tree, sum / N, rsqrtf, +// then w * (rstd * x) and one cast. nn.RMSNorm therefore matches bitwise, and +// each modulation step rounds to the tensor dtype exactly where the eager +// `n * (1.0 + scale) + shift` expression does. Rows are independent: the result +// does not depend on batch size, position or the other rows. +// +// Backward (FP32, one cast per output, no atomics): +// d_n = g * cast(1 + scale) (g without modulation) +// dx = rstd * w * d_n - x * rstd^3 * sum(w * d_n * x) / N (row-local) +// dweight = sum_r d_n * x * rstd over fixed 256-row tiles, tiles folded in order +// dshift = sum_{r -> i} g, dscale = sum_{r -> i} g * n (sorted tiles) +// +// The tile partials and the folds are also bound separately (`*_partials`, +// `*_fold_*`) so a sequence-parallel caller can compute each WS1 tile on the +// rank that holds its rows and fold the gathered partials in WS1 order. + +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +namespace { + +constexpr int kVec = 4; // PyTorch's layer-norm vec_size +constexpr int kWarps = 4; // num_threads() / warp_size = 128 / 32 +constexpr int kRowTile = 256; // dweight row tile + +__device__ __forceinline__ float to_f(float v) { return v; } +__device__ __forceinline__ float to_f(__nv_bfloat16 v) { return __bfloat162float(v); } +__device__ __forceinline__ float to_f(__half v) { return __half2float(v); } + +template +__device__ __forceinline__ T from_f(float v); +template <> +__device__ __forceinline__ float from_f(float v) { return v; } +template <> +__device__ __forceinline__ __nv_bfloat16 from_f<__nv_bfloat16>(float v) { return __float2bfloat16(v); } +template <> +__device__ __forceinline__ __half from_f<__half>(float v) { return __float2half(v); } + +// Round to T and back: the value an eager T-dtype op would produce. +template +__device__ __forceinline__ float round_t(float v) { return to_f(from_f(v)); } + +// PyTorch's rms statistics: sequential per thread, shuffle-down, cross-warp tree. +// Returns sum(v^2) / N on every thread. +template +__device__ __forceinline__ float row_mean_square(const T* __restrict__ row, int n, float* buf) { + const int numx = blockDim.x * blockDim.y; + const int thrx = threadIdx.x + threadIdx.y * blockDim.x; + float s = 0.0f; + for (int i = thrx; i < n / kVec; i += numx) { +#pragma unroll + for (int k = 0; k < kVec; ++k) { + const float v = to_f(row[i * kVec + k]); + s = s + v * v; + } + } + for (int off = 16; off > 0; off >>= 1) s = s + __shfl_down_sync(0xffffffffu, s, off); + for (int off = blockDim.y / 2; off > 0; off /= 2) { + if (threadIdx.x == 0 && threadIdx.y >= off && threadIdx.y < 2 * off) buf[threadIdx.y - off] = s; + __syncthreads(); + if (threadIdx.x == 0 && threadIdx.y < off) s = s + buf[threadIdx.y]; + __syncthreads(); + } + if (thrx == 0) buf[0] = s / float(n); + __syncthreads(); + const float out = buf[0]; + __syncthreads(); + return out; +} + +// Same tree for an arbitrary per-element FP32 term; returns the row sum. +__device__ __forceinline__ float row_sum_tree(float s, float* buf) { + for (int off = 16; off > 0; off >>= 1) s = s + __shfl_down_sync(0xffffffffu, s, off); + for (int off = blockDim.y / 2; off > 0; off /= 2) { + if (threadIdx.x == 0 && threadIdx.y >= off && threadIdx.y < 2 * off) buf[threadIdx.y - off] = s; + __syncthreads(); + if (threadIdx.x == 0 && threadIdx.y < off) s = s + buf[threadIdx.y]; + __syncthreads(); + } + const int thrx = threadIdx.x + threadIdx.y * blockDim.x; + if (thrx == 0) buf[0] = s; + __syncthreads(); + const float out = buf[0]; + __syncthreads(); + return out; +} + +struct Modulation { + const void* shift; // (R, N) rows with stride `row_stride` (elements), or nullptr + const void* scale; + int64_t row_stride; + int64_t num_rows; + const int64_t* index; // (S,) table row per sequence position + int64_t seq; // S: row r of x uses index[r % S] +}; + +template +__global__ void __launch_bounds__(32 * kWarps) + rmsnorm_modulate_fwd_kernel(const T* __restrict__ x, const T* __restrict__ w, T* __restrict__ y, + float* __restrict__ rstd_out, int n, float eps, Modulation mod) { + __shared__ float buf[kWarps]; + const int64_t r = blockIdx.x; + const T* xr = x + r * n; + const float rstd = rsqrtf(row_mean_square(xr, n, buf) + eps); + const T* shift = nullptr; + const T* scale = nullptr; + if (mod.index != nullptr) { + const int64_t i = mod.index[r % mod.seq]; + CUDA_KERNEL_ASSERT(i >= 0 && i < mod.num_rows); + shift = static_cast(mod.shift) + i * mod.row_stride; + scale = static_cast(mod.scale) + i * mod.row_stride; + } + const int numx = blockDim.x * blockDim.y; + const int thrx = threadIdx.x + threadIdx.y * blockDim.x; + for (int i = thrx; i < n / kVec; i += numx) { +#pragma unroll + for (int k = 0; k < kVec; ++k) { + const int j = i * kVec + k; + float v = round_t(__fmul_rn(to_f(w[j]), __fmul_rn(rstd, to_f(xr[j])))); + if (shift != nullptr) { + // Separately rounded ops (no FMA contraction), as the eager expression. + const float t1 = round_t(__fadd_rn(1.0f, to_f(scale[j]))); + v = round_t(__fadd_rn(round_t(__fmul_rn(v, t1)), to_f(shift[j]))); + } + y[r * n + j] = from_f(v); + } + } + if (thrx == 0) rstd_out[r] = rstd; +} + +template +__device__ __forceinline__ float d_norm_out(const T* g_row, const T* scale, int j) { + const float g = to_f(g_row[j]); + return scale == nullptr ? g : g * round_t(1.0f + to_f(scale[j])); +} + +template +__global__ void __launch_bounds__(32 * kWarps) + rmsnorm_modulate_dx_kernel(const T* __restrict__ g, const T* __restrict__ x, + const T* __restrict__ w, const float* __restrict__ rstd, + T* __restrict__ dx, int n, Modulation mod) { + __shared__ float buf[kWarps]; + const int64_t r = blockIdx.x; + const T* xr = x + r * n; + const T* gr = g + r * n; + const T* scale = nullptr; + if (mod.index != nullptr) { + const int64_t i = mod.index[r % mod.seq]; + CUDA_KERNEL_ASSERT(i >= 0 && i < mod.num_rows); + scale = static_cast(mod.scale) + i * mod.row_stride; + } + const int numx = blockDim.x * blockDim.y; + const int thrx = threadIdx.x + threadIdx.y * blockDim.x; + float s = 0.0f; + for (int i = thrx; i < n / kVec; i += numx) { +#pragma unroll + for (int k = 0; k < kVec; ++k) { + const int j = i * kVec + k; + s = fmaf(to_f(w[j]) * d_norm_out(gr, scale, j), to_f(xr[j]), s); + } + } + const float dot = row_sum_tree(s, buf); + const float rs = rstd[r]; + const float coef = rs * rs * rs * dot / float(n); + for (int i = thrx; i < n / kVec; i += numx) { +#pragma unroll + for (int k = 0; k < kVec; ++k) { + const int j = i * kVec + k; + dx[r * n + j] = from_f(rs * to_f(w[j]) * d_norm_out(gr, scale, j) - to_f(xr[j]) * coef); + } + } +} + +// partial[tile, j] = sum over rows of the tile (ascending) of d_n * x * rstd. +// Tiles are the fixed 256-row blocks, or (sequence parallel) explicit row lists +// rows[tile_begin[tile] .. tile_end[tile]) that hold the same rows in the same order. +template +__global__ void rmsnorm_dweight_partial_kernel(const T* __restrict__ g, const T* __restrict__ x, + const float* __restrict__ rstd, + float* __restrict__ partial, int64_t rows, int n, + Modulation mod, + const int64_t* __restrict__ row_list = nullptr, + const int64_t* __restrict__ tile_begin = nullptr, + const int64_t* __restrict__ tile_end = nullptr) { + const int j = blockIdx.y * blockDim.x + threadIdx.x; + if (j >= n) return; + const int64_t tile = blockIdx.x; + const int64_t p0 = row_list != nullptr ? tile_begin[tile] : tile * kRowTile; + const int64_t p1 = row_list != nullptr ? tile_end[tile] : min(p0 + kRowTile, rows); + float acc = 0.0f; + for (int64_t p = p0; p < p1; ++p) { + const int64_t r = row_list != nullptr ? row_list[p] : p; + const T* scale = nullptr; + if (mod.index != nullptr) { + const int64_t i = mod.index[r % mod.seq]; + CUDA_KERNEL_ASSERT(i >= 0 && i < mod.num_rows); + scale = static_cast(mod.scale) + i * mod.row_stride; + } + acc = fmaf(d_norm_out(g + r * n, scale, j), to_f(x[r * n + j]) * rstd[r], acc); + } + partial[static_cast(blockIdx.x) * n + j] = acc; +} + +template +__global__ void fold_tiles_kernel(const float* __restrict__ partial, T* __restrict__ out, + int64_t tiles, int n) { + const int j = blockIdx.x * blockDim.x + threadIdx.x; + if (j >= n) return; + float acc = 0.0f; + for (int64_t t = 0; t < tiles; ++t) acc += partial[t * n + j]; + out[j] = from_f(acc); +} + +// Table gradient tiles: positions sorted by (index, position); columns j < n +// accumulate g (shift), columns j >= n accumulate g * n_out (scale). +template +__global__ void modulation_grad_partial_kernel( + const T* __restrict__ g, const T* __restrict__ x, const T* __restrict__ w, + const float* __restrict__ rstd, const int64_t* __restrict__ sorted_pos, + const int64_t* __restrict__ tile_begin, const int64_t* __restrict__ tile_end, + float* __restrict__ partial, int n) { + const int64_t tile = blockIdx.x; + const int j2 = blockIdx.y * blockDim.x + threadIdx.x; + if (j2 >= 2 * n) return; + const bool is_scale = j2 >= n; + const int j = is_scale ? j2 - n : j2; + const float wj = to_f(w[j]); + float acc = 0.0f; + for (int64_t p = tile_begin[tile]; p < tile_end[tile]; ++p) { + const int64_t r = sorted_pos[p]; + const float gv = to_f(g[r * n + j]); + if (is_scale) { + const float nv = round_t(wj * (rstd[r] * to_f(x[r * n + j]))); + acc = fmaf(gv, nv, acc); + } else { + acc += gv; + } + } + partial[tile * 2 * n + j2] = acc; +} + +__global__ void fold_segments_kernel(const float* __restrict__ partial, + const int64_t* __restrict__ seg_first_tile, + float* __restrict__ out, int width) { + const int64_t seg = blockIdx.x; + const int j = blockIdx.y * blockDim.x + threadIdx.x; + if (j >= width) return; + float acc = 0.0f; + for (int64_t t = seg_first_tile[seg]; t < seg_first_tile[seg + 1]; ++t) acc += partial[t * width + j]; + out[seg * width + j] = acc; +} + +void check_rows(const torch::Tensor& t, const char* name, const torch::Tensor& like) { + TORCH_CHECK(t.is_cuda() && t.device() == like.device(), name, " must be on ", like.device()); + TORCH_CHECK(t.dim() == 2 && t.is_contiguous(), name, " must be a contiguous 2-D tensor"); + TORCH_CHECK(t.scalar_type() == like.scalar_type(), name, " must have dtype ", like.scalar_type()); +} + +Modulation make_modulation(const c10::optional& shift, + const c10::optional& scale, + const c10::optional& index, const torch::Tensor& x, + int64_t n) { + Modulation mod{nullptr, nullptr, 0, 0, nullptr, 1}; + if (!index.has_value()) { + TORCH_CHECK(!shift.has_value() && !scale.has_value(), "shift/scale need a row index"); + return mod; + } + TORCH_CHECK(shift.has_value() && scale.has_value(), "modulation needs both shift and scale"); + const auto& sh = *shift; + const auto& sc = *scale; + const auto& ix = *index; + for (const auto* t : {&sh, &sc}) { + TORCH_CHECK(t->is_cuda() && t->device() == x.device() && t->scalar_type() == x.scalar_type(), + "shift/scale must be CUDA tensors with x's dtype"); + TORCH_CHECK(t->dim() == 2 && t->size(1) == n && t->stride(1) == 1, + "shift/scale must be (R, N) with unit column stride"); + } + TORCH_CHECK(sh.size(0) == sc.size(0) && sh.stride(0) == sc.stride(0), + "shift and scale must be views of the same table layout"); + TORCH_CHECK(ix.is_cuda() && ix.device() == x.device() && ix.scalar_type() == at::kLong && + ix.dim() == 1 && ix.is_contiguous() && ix.numel() > 0, + "row index must be a non-empty contiguous int64 tensor on ", x.device()); + TORCH_CHECK(x.size(0) % ix.size(0) == 0, "x rows (", x.size(0), ") must be a multiple of S (", + ix.size(0), ")"); + mod.shift = sh.data_ptr(); + mod.scale = sc.data_ptr(); + mod.row_stride = sh.stride(0); + mod.num_rows = sh.size(0); + mod.index = ix.data_ptr(); + mod.seq = ix.size(0); + return mod; +} + +void check_xw(const torch::Tensor& x, const torch::Tensor& w) { + TORCH_CHECK(x.is_cuda() && x.dim() == 2 && x.is_contiguous(), "x must be a contiguous (M, N) CUDA tensor"); + TORCH_CHECK(x.size(0) > 0, "x must have at least one row"); + TORCH_CHECK(w.is_cuda() && w.device() == x.device() && w.dim() == 1 && + w.size(0) == x.size(1) && w.is_contiguous() && + w.scalar_type() == x.scalar_type(), + "weight must be a contiguous (N,) CUDA tensor with x's dtype and device"); + TORCH_CHECK(x.size(1) > 0, "x must have at least one column"); + TORCH_CHECK(x.size(1) % kVec == 0, "N=", x.size(1), " must be a multiple of ", kVec, + " (PyTorch's vectorized RMSNorm path)"); +} + +// float32 / float16 / bfloat16 only (float64 has no PyTorch-matching vector path here). +#define H3_DISPATCH(TYPE, NAME, ...) \ + [&] { \ + switch (TYPE) { \ + case at::kFloat: { using T = float; __VA_ARGS__(); break; } \ + case at::kHalf: { using T = __half; __VA_ARGS__(); break; } \ + case at::kBFloat16: { using T = __nv_bfloat16; __VA_ARGS__(); break; } \ + default: TORCH_CHECK(false, NAME, ": unsupported dtype ", TYPE); \ + } \ + }() + +} // namespace + +// x (M, N); optional shift/scale (R, N) strided views and index (S,), M % S == 0. +std::vector h3_rmsnorm_forward(torch::Tensor x, torch::Tensor weight, double eps, + c10::optional shift, + c10::optional scale, + c10::optional index) { + check_xw(x, weight); + TORCH_CHECK(x.scalar_type() != at::kDouble, "float64 is not supported"); + const int64_t n = x.size(1); + const Modulation mod = make_modulation(shift, scale, index, x, n); + const c10::cuda::CUDAGuard guard(x.device()); + auto y = torch::empty_like(x); + auto rstd = torch::empty({x.size(0)}, x.options().dtype(at::kFloat)); + auto stream = at::cuda::getCurrentCUDAStream(); + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_forward", [&] { + rmsnorm_modulate_fwd_kernel<<(x.size(0)), dim3(32, kWarps), 0, stream>>>( + reinterpret_cast(x.data_ptr()), reinterpret_cast(weight.data_ptr()), + reinterpret_cast(y.data_ptr()), rstd.data_ptr(), static_cast(n), + static_cast(eps), mod); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {y, rstd}; +} + +// Returns {dx, dweight} and, with modulation, {d_shift, d_scale} as FP32 (R, N). +std::vector h3_rmsnorm_backward( + torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, + c10::optional shift, c10::optional scale, + c10::optional index, c10::optional sorted_pos, + c10::optional tile_begin, c10::optional tile_end, + c10::optional seg_first_tile) { + check_xw(x, weight); + check_rows(grad, "grad", x); + TORCH_CHECK(grad.sizes() == x.sizes(), "grad must match x"); + TORCH_CHECK(rstd.is_cuda() && rstd.device() == x.device() && + rstd.scalar_type() == at::kFloat && rstd.dim() == 1 && + rstd.is_contiguous() && rstd.size(0) == x.size(0), + "rstd must be contiguous float32 (M,) statistics on x's device"); + const int64_t rows = x.size(0); + const int64_t n = x.size(1); + const Modulation mod = make_modulation(shift, scale, index, x, n); + if (mod.index != nullptr) { + TORCH_CHECK(sorted_pos.has_value() && tile_begin.has_value() && tile_end.has_value() && + seg_first_tile.has_value(), + "modulated backward needs the sorted segment tiles"); + const std::pair tiles[] = {{&*sorted_pos, "sorted_pos"}, + {&*tile_begin, "tile_begin"}, + {&*tile_end, "tile_end"}, + {&*seg_first_tile, "seg_first_tile"}}; + for (const auto& [t, name] : tiles) { + TORCH_CHECK(t->is_cuda() && t->device() == x.device() && t->scalar_type() == at::kLong && + t->dim() == 1 && t->is_contiguous(), + name, ": tile metadata must be contiguous int64 tensors on ", x.device()); + } + TORCH_CHECK(sorted_pos->size(0) == rows, "sorted_pos must have M entries"); + TORCH_CHECK(seg_first_tile->size(0) == shift->size(0) + 1, + "seg_first_tile must have R + 1 entries"); + TORCH_CHECK(tile_begin->size(0) == tile_end->size(0), + "tile_begin and tile_end must have the same number of entries"); + } else { + TORCH_CHECK(!sorted_pos.has_value() && !tile_begin.has_value() && !tile_end.has_value() && + !seg_first_tile.has_value(), + "sorted segment tiles require modulation"); + } + const c10::cuda::CUDAGuard guard(x.device()); + auto stream = at::cuda::getCurrentCUDAStream(); + auto dx = torch::empty_like(x); + const int64_t tiles = (rows + kRowTile - 1) / kRowTile; + auto partial = torch::empty({tiles, n}, x.options().dtype(at::kFloat)); + auto dweight = torch::empty_like(weight); + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_backward", [&] { + const T* g = reinterpret_cast(grad.data_ptr()); + const T* xp = reinterpret_cast(x.data_ptr()); + const T* wp = reinterpret_cast(weight.data_ptr()); + rmsnorm_modulate_dx_kernel<<(rows), dim3(32, kWarps), 0, stream>>>( + g, xp, wp, rstd.data_ptr(), reinterpret_cast(dx.data_ptr()), + static_cast(n), mod); + rmsnorm_dweight_partial_kernel<<(tiles), col_blocks), threads, 0, + stream>>>(g, xp, rstd.data_ptr(), + partial.data_ptr(), rows, + static_cast(n), mod); + fold_tiles_kernel<<>>( + partial.data_ptr(), reinterpret_cast(dweight.data_ptr()), tiles, + static_cast(n)); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + if (mod.index == nullptr) return {dx, dweight}; + + const int64_t segments = seg_first_tile->numel() - 1; + const int64_t seg_tiles = tile_begin->numel(); + const int64_t width = 2 * n; + auto seg_partial = torch::empty({std::max(seg_tiles, 1), width}, x.options().dtype(at::kFloat)); + auto table_grad = torch::empty({segments, width}, x.options().dtype(at::kFloat)); + const unsigned wide_blocks = static_cast((width + threads - 1) / threads); + if (seg_tiles > 0) { + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_modulation_grad", [&] { + modulation_grad_partial_kernel + <<(seg_tiles), wide_blocks), threads, 0, stream>>>( + reinterpret_cast(grad.data_ptr()), reinterpret_cast(x.data_ptr()), + reinterpret_cast(weight.data_ptr()), rstd.data_ptr(), + sorted_pos->data_ptr(), tile_begin->data_ptr(), + tile_end->data_ptr(), seg_partial.data_ptr(), static_cast(n)); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + fold_segments_kernel<<(segments), wide_blocks), threads, 0, stream>>>( + seg_partial.data_ptr(), seg_first_tile->data_ptr(), + table_grad.data_ptr(), static_cast(width)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {dx, dweight, table_grad.narrow(1, 0, n), table_grad.narrow(1, n, n)}; +} + +namespace { + +void check_tiles(const torch::Tensor& rows, const torch::Tensor& begin, const torch::Tensor& end, + const torch::Tensor& like) { + for (const auto* t : {&rows, &begin, &end}) { + TORCH_CHECK(t->is_cuda() && t->device() == like.device() && t->scalar_type() == at::kLong && + t->dim() == 1 && t->is_contiguous(), + "tile metadata must be contiguous int64 CUDA tensors on x's device"); + } + TORCH_CHECK(begin.numel() == end.numel(), "tile_begin and tile_end differ in length"); +} + +} // namespace + +// Sequence parallel: the dweight and (with modulation) table-gradient tile +// partials of the WS1 backward, for explicit tiles over the rows of x. `index` +// gives each row's table row (one entry per row of x). Returns +// {dweight_partial (tiles, N)} or {dweight_partial, table_partial (seg_tiles, 2N)}. +std::vector h3_rmsnorm_backward_partials( + torch::Tensor grad, torch::Tensor x, torch::Tensor weight, torch::Tensor rstd, + c10::optional shift, c10::optional scale, + c10::optional index, torch::Tensor dw_rows, torch::Tensor dw_begin, + torch::Tensor dw_end, c10::optional seg_rows, + c10::optional seg_begin, c10::optional seg_end) { + check_xw(x, weight); + check_rows(grad, "grad", x); + TORCH_CHECK(grad.sizes() == x.sizes(), "grad must match x"); + TORCH_CHECK(rstd.is_cuda() && rstd.device() == x.device() && rstd.scalar_type() == at::kFloat && + rstd.numel() == x.size(0), + "rstd must be the forward's float32 (M,) statistics on x's device"); + check_tiles(dw_rows, dw_begin, dw_end, x); + const int64_t rows = x.size(0); + const int64_t n = x.size(1); + const Modulation mod = make_modulation(shift, scale, index, x, n); + TORCH_CHECK(mod.index == nullptr || mod.seq == rows, "index needs one entry per row of x"); + const c10::cuda::CUDAGuard guard(x.device()); + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + const int64_t tiles = dw_begin.numel(); + auto dw_partial = torch::empty({tiles, n}, x.options().dtype(at::kFloat)); + if (tiles > 0) { + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_backward_partials", [&] { + rmsnorm_dweight_partial_kernel<<(tiles), col_blocks), threads, + 0, stream>>>( + reinterpret_cast(grad.data_ptr()), reinterpret_cast(x.data_ptr()), + rstd.data_ptr(), dw_partial.data_ptr(), rows, static_cast(n), mod, + dw_rows.data_ptr(), dw_begin.data_ptr(), dw_end.data_ptr()); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + if (mod.index == nullptr) return {dw_partial}; + + TORCH_CHECK(seg_rows.has_value() && seg_begin.has_value() && seg_end.has_value(), + "modulated partials need the segment tiles"); + check_tiles(*seg_rows, *seg_begin, *seg_end, x); + const int64_t seg_tiles = seg_begin->numel(); + const int64_t width = 2 * n; + auto seg_partial = torch::empty({seg_tiles, width}, x.options().dtype(at::kFloat)); + const unsigned wide_blocks = static_cast((width + threads - 1) / threads); + if (seg_tiles > 0) { + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_modulation_partials", [&] { + modulation_grad_partial_kernel + <<(seg_tiles), wide_blocks), threads, 0, stream>>>( + reinterpret_cast(grad.data_ptr()), reinterpret_cast(x.data_ptr()), + reinterpret_cast(weight.data_ptr()), rstd.data_ptr(), + seg_rows->data_ptr(), seg_begin->data_ptr(), + seg_end->data_ptr(), seg_partial.data_ptr(), static_cast(n)); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + } + return {dw_partial, seg_partial}; +} + +// The WS1 folds: dweight = ascending fold of all tiles, cast to weight's dtype; +// table rows = ascending fold of each segment's tiles, FP32 (R, 2N). +std::vector h3_rmsnorm_fold_partials(torch::Tensor dw_partial, torch::Tensor weight, + c10::optional seg_partial, + c10::optional seg_first_tile) { + TORCH_CHECK(dw_partial.is_cuda() && dw_partial.dim() == 2 && dw_partial.is_contiguous() && + dw_partial.scalar_type() == at::kFloat, + "dw_partial must be a contiguous float32 (tiles, N) CUDA tensor"); + TORCH_CHECK(weight.is_cuda() && weight.device() == dw_partial.device() && weight.dim() == 1 && + weight.size(0) == dw_partial.size(1), + "weight must be (N,) on dw_partial's device"); + const int64_t n = dw_partial.size(1); + const c10::cuda::CUDAGuard guard(dw_partial.device()); + auto stream = at::cuda::getCurrentCUDAStream(); + const int threads = 256; + const unsigned col_blocks = static_cast((n + threads - 1) / threads); + auto dweight = torch::empty_like(weight); + H3_DISPATCH(weight.scalar_type(), "h3_rmsnorm_fold_partials", [&] { + fold_tiles_kernel<<>>( + dw_partial.data_ptr(), reinterpret_cast(dweight.data_ptr()), + dw_partial.size(0), static_cast(n)); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + if (!seg_partial.has_value()) return {dweight}; + + TORCH_CHECK(seg_first_tile.has_value() && seg_first_tile->is_cuda() && + seg_first_tile->device() == dw_partial.device() && + seg_first_tile->scalar_type() == at::kLong && seg_first_tile->is_contiguous(), + "segment fold needs seg_first_tile (int64, on dw_partial's device)"); + TORCH_CHECK(seg_partial->is_cuda() && seg_partial->device() == dw_partial.device() && + seg_partial->dim() == 2 && seg_partial->is_contiguous() && + seg_partial->scalar_type() == at::kFloat && seg_partial->size(1) == 2 * n, + "seg_partial must be a contiguous float32 (tiles, 2N) CUDA tensor"); + const int64_t segments = seg_first_tile->numel() - 1; + const int64_t width = 2 * n; + auto table_grad = torch::empty({segments, width}, dw_partial.options()); + const unsigned wide_blocks = static_cast((width + threads - 1) / threads); + fold_segments_kernel<<(segments), wide_blocks), threads, 0, stream>>>( + seg_partial->data_ptr(), seg_first_tile->data_ptr(), + table_grad.data_ptr(), static_cast(width)); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return {dweight, table_grad.narrow(1, 0, n), table_grad.narrow(1, n, n)}; +} + +// Row-local part of the WS1 backward only: dx (sequence-parallel callers +// compute the cross-row reductions with h3_rmsnorm_backward_partials). +torch::Tensor h3_rmsnorm_backward_dx(torch::Tensor grad, torch::Tensor x, torch::Tensor weight, + torch::Tensor rstd, c10::optional shift, + c10::optional scale, + c10::optional index) { + check_xw(x, weight); + check_rows(grad, "grad", x); + TORCH_CHECK(grad.sizes() == x.sizes(), "grad must match x"); + TORCH_CHECK(rstd.is_cuda() && rstd.device() == x.device() && rstd.scalar_type() == at::kFloat && + rstd.numel() == x.size(0), + "rstd must be the forward's float32 (M,) statistics on x's device"); + const Modulation mod = make_modulation(shift, scale, index, x, x.size(1)); + const c10::cuda::CUDAGuard guard(x.device()); + auto dx = torch::empty_like(x); + H3_DISPATCH(x.scalar_type(), "h3_rmsnorm_backward_dx", [&] { + rmsnorm_modulate_dx_kernel<<(x.size(0)), dim3(32, kWarps), 0, + at::cuda::getCurrentCUDAStream()>>>( + reinterpret_cast(grad.data_ptr()), reinterpret_cast(x.data_ptr()), + reinterpret_cast(weight.data_ptr()), rstd.data_ptr(), + reinterpret_cast(dx.data_ptr()), static_cast(x.size(1)), mod); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return dx; +} diff --git a/csrc/cuda/h3/timestep_sinusoid.cu b/csrc/cuda/h3/timestep_sinusoid.cu new file mode 100644 index 000000000..9322920a4 --- /dev/null +++ b/csrc/cuda/h3/timestep_sinusoid.cu @@ -0,0 +1,90 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 RL-Kernel Contributors +// +// MiniMax-H3 sinusoidal timestep features (RFC #420 `timestep_sinusoid_h3`). +// +// freq[k] = expf((c * k) * (1 / half)) c = (float)(-ln(max_period)) +// arg[t, k] = t[t] * freq[k] +// out[t] = [cosf(arg) | sinf(arg)] (T, 2 * half) FP32 +// +// Every value is produced by the same FP32 operation sequence as diffusers' +// get_timestep_embedding on CUDA (scalar * arange, multiply by the reciprocal +// of the CPU-scalar divisor, exp, multiply, cos/sin), so the output is meant +// to be bitwise equal to that path. Each element depends only on (t, k): +// batch-, position- and repeat-invariant by construction. Build without +// --use_fast_math: expf/sinf/cosf must stay the precise libdevice functions. + +#include +#include +#include +#include + +#include +#include + +namespace { + +__global__ void h3_timestep_sinusoid_kernel( + const float* __restrict__ timestep, + float* __restrict__ out, + int64_t num_timesteps, + int64_t half, + float neg_log_max_period, + float inv_half) { + const int64_t total = num_timesteps * half; + for (int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; idx < total; + idx += static_cast(gridDim.x) * blockDim.x) { + const int64_t row = idx / half; + const int64_t k = idx - row * half; + // __fmul_rn keeps each product a separately rounded FP32 multiply (no FMA + // contraction), matching the two elementwise torch kernels it replays. + const float exponent = __fmul_rn(__fmul_rn(neg_log_max_period, static_cast(k)), inv_half); + const float freq = expf(exponent); + const float arg = __fmul_rn(timestep[row], freq); + float* out_row = out + row * 2 * half; + out_row[k] = cosf(arg); + out_row[half + k] = sinf(arg); + } +} + +} // namespace + +torch::Tensor h3_timestep_sinusoid_forward(torch::Tensor timestep, int64_t num_channels, + double max_period, bool check_range) { + TORCH_CHECK(timestep.is_cuda(), "timestep must be a CUDA tensor"); + TORCH_CHECK(timestep.scalar_type() == at::kFloat, "timestep must be float32, got ", + timestep.scalar_type()); + TORCH_CHECK(timestep.dim() == 1, "timestep must be 1-D (num_timesteps,)"); + TORCH_CHECK(timestep.numel() > 0, "timestep must hold at least one timestep"); + TORCH_CHECK(num_channels > 0 && num_channels % 2 == 0, + "num_channels must be a positive even number, got ", num_channels); + TORCH_CHECK(max_period > 0.0, "max_period must be positive"); + + const c10::cuda::CUDAGuard device_guard(timestep.device()); + auto t = timestep.contiguous(); + // Check at the native boundary, including direct extension calls. The explicit + // opt-out is for already-validated inputs and kernel-only profiling. + if (check_range) { + TORCH_CHECK_VALUE(((t >= 0) & (t <= 1)).all().item(), + "timestep must be finite and lie in [0, 1]: H3 consumes t = 1 - sigma " + "unscaled"); + } + const int64_t num_timesteps = t.size(0); + const int64_t half = num_channels / 2; + auto out = torch::empty({num_timesteps, num_channels}, t.options()); + + // torch evaluates `-math.log(max_period) * arange` with the Python double + // rounded to the FP32 opmath type, and `x / half` as x * (1 / half) in FP32. + const float neg_log_max_period = static_cast(-std::log(max_period)); + const float inv_half = 1.0f / static_cast(half); + + const int threads = 256; + const int64_t total = num_timesteps * half; + const int64_t blocks = std::min((total + threads - 1) / threads, 65535); + auto stream = at::cuda::getCurrentCUDAStream(); + h3_timestep_sinusoid_kernel<<(blocks), threads, 0, stream>>>( + t.data_ptr(), out.data_ptr(), num_timesteps, half, neg_log_max_period, + inv_half); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} diff --git a/docs/.nav.yml b/docs/.nav.yml index 29b5ead37..39e42f7fe 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -29,6 +29,16 @@ nav: - operators/sampling.md - operators/det-gemm.md - operators/embedding.md + - MiniMax-H3: + - operators/h3-timestep-sinusoid.md + - operators/h3-timestep-mlp.md + - operators/h3-adaln-projection.md + - operators/h3-adaln-row-gather.md + - operators/h3-rmsnorm.md + - operators/h3-adaln-gate-residual.md + - operators/h3-final-adaln-out.md + - operators/h3-tp-adaln-3mod.md + - operators/h3-sp-norm-adaln.md - Developer Guide: - contributing/README.md - Contributor Guide: contributing/contributor-guide.md diff --git a/docs/operators/README.md b/docs/operators/README.md index 00f4cbb45..c99451553 100644 --- a/docs/operators/README.md +++ b/docs/operators/README.md @@ -32,4 +32,14 @@ Every operator page should include: - [Matmul](matmul.md) - [Sampling](sampling.md) - [Token Embedding](embedding.md) +- [MiniMax-H3 Timestep Sinusoid](h3-timestep-sinusoid.md) +- [MiniMax-H3 FP32 Timestep MLP](h3-timestep-mlp.md) +- [MiniMax-H3 AdaLN Projection](h3-adaln-projection.md) +- [MiniMax-H3 AdaLN Row Gather](h3-adaln-row-gather.md) +- [MiniMax-H3 RMSNorm and AdaLN Modulation](h3-rmsnorm.md) +- [MiniMax-H3 AdaLN Gated Residual](h3-adaln-gate-residual.md) +- [MiniMax-H3 Final AdaLN Output](h3-final-adaln-out.md) +- [MiniMax-H3 Tensor-Parallel AdaLN Projection](h3-tp-adaln-3mod.md) +- [MiniMax-H3 Sequence-Parallel Norm and AdaLN](h3-sp-norm-adaln.md) +- [MiniMax-H3 One Block, Forward and Backward](h3-one-block.md) - [Operator Doc Template](../contributing/operator-doc-template.md) diff --git a/docs/operators/h3-adaln-gate-residual.md b/docs/operators/h3-adaln-gate-residual.md new file mode 100644 index 000000000..dc6c7c702 --- /dev/null +++ b/docs/operators/h3-adaln-gate-residual.md @@ -0,0 +1,124 @@ +# MiniMax-H3 AdaLN Gated Residual + +## Summary + +`adaln_gate_residual` adds each sublayer's output back to the residual stream through the +per-row AdaLN gate (RFC #420, WS1 step 6). In every H3 block this happens twice, after attention +with `gate_msa` and after the FFN with `gate_mlp`: + +```text +hidden = residual + gate[adaln_indices] * sublayer_output # RFC #420 §4 order +``` + +`gate` is an `(R, H)` row view of the [AdaLN projection](h3-adaln-projection.md) table. The op +gathers the gate row inside the kernel, so the gathered `(S, H)` gate is never materialised. + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("adaln_gate_residual", device="cuda") +hidden = op(residual, attn_output, gate_msa, adaln_indices) +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA | `rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual.H3GateResidualCudaOp` | `rl_engine._C.h3_gate_residual_{forward,backward}` | Forward and `d_sublayer` bitwise equal to diffusers | +| PyTorch reference | `rl_engine.reference.minimax_h3.gate_residual.NativeH3GateResidualOp` | n/a | Eager forward with a deterministic gate gradient; `forward_fp32`: FP64 golden | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `residual`, `sublayer_output` | `(..., S, H)` | bf16 / fp16 / fp32 | Same shape and dtype | +| `gate` | `(R, H)` row view | same | Unit column stride | +| `index` | `(S,)` | int64 | In `[0, R)`; shared across the batch | + +## Numerics + +- **Forward.** Each element computes `p = round(gate * y)` and then `out = round(residual + p)`. + `__fmul_rn`/`__fadd_rn` keep FP32 from contracting the two operations into one FMA. The + roundings fall exactly where the eager expression rounds, so the forward is **bitwise equal to + diffusers** in BF16, FP16 and FP32. Rows are independent. +- **Backward.** + - `d_residual` is the incoming gradient. + - `d_sublayer = round(grad * gate)`. The product of two 16-bit values is exact in FP32, so this + is bitwise equal to the eager VJP. + - `d_gate` is an FP32 segmented sum of `grad * sublayer_output` over positions sorted stably + by table row, using the same tiles as [`adaln_row_gather`](h3-adaln-row-gather.md), rounded + once. + + Diffusers' `d_gate` uses `index_select`'s BF16 atomic scatter-add instead. It is + non-deterministic and 13–49× further from the FP64 golden (its error varies from run to run). + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op adaln_gate_residual +``` + +B200, B = 1, H = 5376, BF16. Backward timings exclude input and forward-graph setup: + +| S | CUDA fwd | diffusers fwd | CUDA bwd | diffusers bwd | +| --- | --- | --- | --- | --- | +| 4097 | 0.07 ms | 0.09 ms | 1.12 ms | 0.58 ms | +| 32768 | 0.24 ms | 0.42 ms | 1.12 ms | 2.24 ms | +| 131072 | 0.88 ms | 1.61 ms | 2.22 ms | 8.92 ms | + +The forward is one pass over `residual` and `sublayer_output`, using 16-byte vectors with one +block per row. At small S the backward pays a fixed cost of about 1 ms for the stable sort and +tile setup. + +## Evidence + +![adaln_gate_residual on B200: latency and backward accuracy](../../reports/experiments/h3-adaln-gate-residual-b200/figure.png) + +The data is in [`report.json`](../../reports/experiments/h3-adaln-gate-residual-b200/report.json), +written by `tools/validation/models/h3_evidence.py` from a clean tree at commit `bde1d8d`. It also records +forward bitwise equality with diffusers in bf16, fp16 and fp32, and row invariance. + +## Existing implementations (RFC #420 reuse rule) + +![gate_residual vs existing implementations](../../reports/experiments/h3-prior-art-b200/gate_residual.png) + +| Implementation | Batch-invariant | size 4097: fwd err / worst grad err / fwd+bwd | size 32768: fwd err / worst grad err / fwd+bwd | +|---|---|---|---| +| diffusers residual + gate.index_select(...) * y | **no** (param/table grads not repeatable) | 3.0e-03 / 4.3e-02 / 637 µs | 2.9e-03 / 1.4e-01 / 2502 µs | +| rl-kernel H3GateResidualCudaOp | yes | 3.0e-03 / 2.6e-03 / 1210 µs | 2.9e-03 / 2.6e-03 / 1596 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +131072-row batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`gate_residual.json`](../../reports/experiments/h3-prior-art-b200/gate_residual.json) +was written from a clean tree at `4010854` by + +```bash +python tools/validation/models/h3_prior_art.py --op gate_residual --out reports/experiments/h3-prior-art-b200/gate_residual.json +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/gate_residual.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_adaln_gate_residual.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end, incl. the gated residual +python tools/validation/operators/check_operator.py --op adaln_gate_residual --candidate cuda --device cuda \ + --dtype bf16 --batch 3 --seq 1365 --normalized-dim 5376 --check-grad +python tools/validation/models/h3_evidence.py --op adaln_gate_residual \ + --out reports/experiments/h3-adaln-gate-residual-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-adaln-gate-residual-b200/report.json +``` + +## Known Limitations + +- Until the attention and FFN rows land, the end-to-end chain feeds a seeded stand-in for the + sublayer output. The op itself is exercised on random and pinned-table inputs. +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. diff --git a/docs/operators/h3-adaln-projection.md b/docs/operators/h3-adaln-projection.md new file mode 100644 index 000000000..76e84f5dd --- /dev/null +++ b/docs/operators/h3-adaln-projection.md @@ -0,0 +1,186 @@ +# MiniMax-H3 Three-Modality AdaLN Projection + +## Summary + +`adaln_projection_3mod` is the per-block `MiniMaxH3AdaLayerNormModulation` +(RFC #420, WS1 step 3). It turns the shared FP32 timestep embedding into the six +modulation tensors of one transformer block, for all three modalities: + +```text +act = silu(temb).to(bf16) SiLU in FP32, one declared cast +table = act @ W.T + b (T, 96768) BF16, W: (96768, 2688) +rows = table.view(3T, 32256) row t * 3 + m, m: 0 video, 1 text, 2 audio +shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = rows.chunk(6, -1) # (3T, 5376) each +``` + +Projection output channel `o = m * 32256 + c * 5376 + h` is chunk `c`, hidden index `h` +of modality `m`. The table rows are laid out `[t0m0, t0m1, t0m2, t1m0, ...]`. That is the +layout `adaln_row_gather` addresses with `timestep_indices * 3 + token_tags`. + +Pinned model: `MiniMaxAI/MiniMax-H3@42ed227`, `transformer_blocks.0.adaln_proj.linear` +(BF16 weight and bias). All 50 blocks have the same shape. + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("adaln_projection_3mod", device="cuda") +shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = op(temb, weight, bias) +table = op.forward_table(temb, weight, bias) # raw (T, 96768) for a fused gather +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA (SM80+, validated on SM100) | `rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection.H3AdaLNProjectionCudaOp` | `rl_engine._C.h3_det_linear_*` | BF16 weights: contract `h3-det-linear-bf16-mma-v1`; FP32 weights: `h3-det-linear-v1` | +| PyTorch reference | `rl_engine.reference.minimax_h3.adaln_projection.NativeH3AdaLNProjectionOp` | n/a | `forward`: provider path; `forward_fp32`: declared-cast FP64 golden | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `temb` | `(T, D)`, `T >= 1` | float32 | H3: `D = 2688`. BF16 is rejected (RFC probe H7) | +| `weight` | `(6 * H * 3, D)` | bfloat16 (checkpoint) or float32 | First dim must be a multiple of 18. H3: `H = 5376` | +| `bias` | `(6 * H * 3,)` | same as `weight` | | +| outputs | six `(3T, H)` views | weight dtype | Views of one `(T, 18H)` table, as in diffusers | + +The CUDA kernel needs `D` to be a multiple of 8 for BF16 or 4 for FP32, and BF16 needs SM80+. These inputs fail +closed: a BF16 `temb`, mismatched dtypes or shapes, an empty `temb`, and mixed devices. + +## Numerics + +- **Mixed-precision boundary.** SiLU runs in FP32 at `temb`'s precision, and the result is + rounded to BF16 exactly once. These are the provider's own elementwise ops, so the + activation is bitwise equal to diffusers. Casting before the SiLU (probe H7) is + rejected at the API. It is also numerically visible: on block 0 it changes about 52% of + the BF16 outputs. +- **Projection (BF16 weights).** It follows contract `h3-det-linear-bf16-mma-v1` in + `csrc/cuda/h3/det_linear.cu`: + - Each warp owns 16 output columns and 8 input rows. Rows past T are zero, so every + T <= 8 runs the same instruction sequence. A tensor-core column never sees another + column's data. + - K is visited in groups of 16 in ascending order. Each `mma.sync m16n8k16` (BF16 in, + FP32 out) starts from zero, so the tensor core only sums 16 products. Its result is + added to an FP32 running sum with an IEEE add. + - Then the bias is added and the result is rounded to BF16 once. + + A timestep's 18 modulation rows are bitwise independent of the other timesteps in the + call and of their order. Restarting the accumulator for every 16-wide group is what keeps + accuracy at the level of an FP32 FMA chain. Letting the tensor core accumulate across + the whole of K, as cuBLAS does, loses about 0.15% of correctly rounded outputs. FP32 + weights use the FMA path of `h3-det-linear-v1` instead, because FP32 on tensor cores + would mean TF32, which the contract forbids. +- **Golden.** `forward_fp32` computes the SiLU in FP64, rounds it to BF16 at the declared + boundary (the cast is model semantics), and runs the projection in FP64. It returns + FP32 without the final rounding. Its gradient stays in FP64 end to end, because the + cast is applied straight-through (identity VJP). Writing the cast as + `.to(bf16).double()` would make autograd round the golden's own `d_temb` to BF16. +- **Backward.** No cuBLAS and no atomics. `d_act` is computed with the chunked + deterministic `grad @ W`, kept in FP32 through the cast (whose VJP is the identity), + and passed through the FP32 SiLU VJP; it is row-local and batch-invariant. + `dW` and `db` are ascending-row FP32 folds, rounded to BF16 once. Unlike the provider, + `d_temb` is not rounded to BF16 on its way through the cast. + +Measured on a B200 with the pinned block-0 weights and torch 2.13.0+cu130. Fractions are +medians over 20 random draws of `temb`, with T = 4: + +| Comparison | CUDA | provider (cuBLAS) | +| --- | --- | --- | +| outputs equal to the correctly rounded golden | 99.982% (worst draw 99.976%) | 99.842% (worst draw 99.787%) | +| CUDA equal to provider | 99.84% | | +| early-cast golden (probe H7) equal to CUDA | 43% | | +| `d_temb` vs FP64 golden (gtest, T = 3) | max abs 2.3e-5 | 6.1e-2 (`d_act` rounded to BF16) | +| `dW`, `db` vs FP64 golden rounded to BF16 | > 99.9% bitwise equal | | + +The contract tolerance is `reduction` / `bfloat16`: atol 5e-2, rtol 2e-2, and atol 1e-1 +for gradients. + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op adaln_projection_3mod +``` + +B200, pinned BF16 weights (520 MB per block, streamed once per call): + +| T | CUDA op | provider (SiLU + cast + cuBLAS) | CUDA GEMV kernel | cuBLAS kernel | +| --- | --- | --- | --- | --- | +| 1 | 105 µs | 115 µs | 78.3 µs (6.6 TB/s) | 93 µs | +| 2 | 105 µs | 103 µs | 78.4 µs | 79 µs | +| 4 | 106 µs | 104 µs | 79.0 µs | 79 µs | + +The kernel takes the same time for every T <= 8, because it always computes 8 rows. It +matches cuBLAS for T >= 2 and is faster at T = 1. The remaining 24 µs of the op time is +the SiLU and cast kernels plus the Python wrapper. With those included, the op is within +about 2% of the provider for T >= 2. + +An earlier FMA-only BF16 kernel ran at 81, 98 and 166 µs for T = 1, 2 and 4. Its time +grew with T because every extra row adds FMAs per weight element. Packed FP32 FMA +(`__ffma2_rn`) gave identical bits but no speedup. Moving to tensor cores removed the +dependence on T. + +## Evidence + +![adaln_projection_3mod on B200: latency and correctly rounded outputs](../../reports/experiments/h3-adaln-projection-b200/figure.png) + +The data is in [`report.json`](../../reports/experiments/h3-adaln-projection-b200/report.json), +written by `tools/validation/models/h3_evidence.py` from a clean tree at commit `d06e120`. The report also +records that a timestep's 18 modulation rows are bitwise identical whether it runs alone +or in a batch of 9. + +## Existing implementations (RFC #420 reuse rule) + +![adaln_projection vs existing implementations](../../reports/experiments/h3-prior-art-b200/adaln_projection.png) + +The sizes below are timestep (input row) counts T. `H3AdaLNProjectionCudaOp` +streams its 520 MB weight once per 8 rows; T = 256 and 2048 require additional +weight passes and are outside its intended T ≤ 8 range. This accounts for the +steep increase in the reported CUDA forward + backward times (77.4 and 627 ms). + +| Implementation | Batch-invariant | size 3: fwd err / worst grad err / fwd+bwd | size 256: fwd err / worst grad err / fwd+bwd | size 2048: fwd err / worst grad err / fwd+bwd | +|---|---|---|---|---| +| diffusers AdaLN projection, BF16 F.linear [plain] | **no** (14214 rows, 3080 sub-batches, 2232 sweep cases) | 3.6e-03 / 3.9e-03 / 716 µs | 3.4e-03 / 3.2e-03 / 728 µs | 3.2e-03 / 3.5e-03 / 2346 µs | +| rl-kernel H3AdaLNProjectionCudaOp | yes | 3.6e-03 / 3.9e-03 / 2276 µs | 3.4e-03 / 3.2e-03 / 77446 µs | 3.2e-03 / 3.5e-03 / 626992 µs | +| diffusers AdaLN projection, BF16 F.linear [vllm] | yes | 3.6e-03 / 3.9e-03 / 681 µs | 3.4e-03 / 3.2e-03 / 722 µs | 3.2e-03 / 3.5e-03 / 2330 µs | +| diffusers AdaLN projection, BF16 F.linear [sglang] | yes | 3.6e-03 / 3.9e-03 / 3096 µs | 3.4e-03 / 3.2e-03 / 3690 µs | 3.2e-03 / 3.5e-03 / 14296 µs | +| diffusers AdaLN projection, BF16 F.linear [sglang_ieee] | yes | 3.6e-03 / 3.9e-03 / 3138 µs | 3.4e-03 / 3.2e-03 / 3740 µs | 3.2e-03 / 3.5e-03 / 14325 µs | +| diffusers AdaLN projection, BF16 F.linear [megatron_te_native] | **no** (7107 rows, 2056 sub-batches, 1104 sweep cases) | 3.6e-03 / 3.9e-03 / 702 µs | 3.4e-03 / 3.2e-03 / 706 µs | 3.2e-03 / 3.5e-03 / 2330 µs | +| diffusers AdaLN projection, BF16 F.linear [megatron_triton] | yes | 3.6e-03 / 3.9e-03 / 2532 µs | 3.4e-03 / 3.2e-03 / 3204 µs | 3.2e-03 / 3.5e-03 / 13771 µs | +| diffusers AdaLN projection, BF16 F.linear [megatron_triton_ieee] | yes | 3.6e-03 / 3.9e-03 / 2533 µs | 3.4e-03 / 3.2e-03 / 3206 µs | 3.2e-03 / 3.5e-03 / 13773 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +4096-timestep batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`adaln_projection.json`](../../reports/experiments/h3-prior-art-b200/adaln_projection.json) +was written from a clean tree at `c88691c` by + +```bash +python tools/validation/models/h3_prior_art.py --op adaln_projection --out reports/experiments/h3-prior-art-b200/adaln_projection.json --megatron-src +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/adaln_projection.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_adaln_projection.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end: sinusoid -> MLP -> projection +python tools/validation/operators/check_operator.py --op adaln_projection_3mod --candidate cuda --device cuda \ + --dtype bf16 --batch 3 --normalized-dim 5376 --check-grad +python tools/validation/models/h3_evidence.py --op adaln_projection_3mod \ + --out reports/experiments/h3-adaln-projection-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-adaln-projection-b200/report.json +``` + +## Known Limitations + +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. +- The BF16 path needs SM80 or newer (`mma.sync`). It is rejected at launch on older GPUs. +- T > 8 runs one more weight pass per 8 rows. +- Sharding the 96768-wide projection (`tp_adaln_3mod`) is a separate WS2 row. diff --git a/docs/operators/h3-adaln-row-gather.md b/docs/operators/h3-adaln-row-gather.md new file mode 100644 index 000000000..ec6772fff --- /dev/null +++ b/docs/operators/h3-adaln-row-gather.md @@ -0,0 +1,181 @@ +# MiniMax-H3 AdaLN Row Gather + +## Summary + +`adaln_row_gather` selects every packed-sequence position's six modulation vectors from +one block's AdaLN table (RFC #420, WS1 step 3): + +```text +adaln_indices = timestep_indices * 3 + token_tags (S,) 0 video, 1 text, 2 audio +shift_msa[adaln_indices], scale_msa[...], gate_msa[...], +shift_mlp[...], scale_mlp[...], gate_mlp[...] six (S, 5376) +``` + +`rows` is the `(3T, 6H)` view of the [AdaLN projection](h3-adaln-projection.md) table, and +its six column blocks are the six modulation tensors. `timestep_indices` and `token_tags` +are semantic inputs (RFC #420 §4): they are validated, never clamped. + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("adaln_row_gather", device="cuda") +outs = op(table.view(-1, 6 * 5376), timestep_indices, token_tags) # six (S, 5376) +outs = op.gather_chunks((shift_msa, ..., gate_mlp), timestep_indices, token_tags) # drop-in +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA (SM90, SM100) | `rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather.H3AdaLNRowGatherCudaOp` | `rl_engine._C.h3_adaln_row_gather_{forward,backward}` | Forward bitwise equal to `index_select`; deterministic backward | +| PyTorch reference | `rl_engine.reference.minimax_h3.adaln_row_gather.NativeH3AdaLNRowGatherOp` | n/a | `index_select` forward; deterministic per-row FP32 backward | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +The raw diffusers path, including its atomic backward, is kept as +`rl_engine.validation.models.h3_provider.provider_adaln_row_gather` for evidence. + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `rows` | `(3T, 6H)` | bf16 / fp16 / fp32 | Unit column stride. Any row stride (16-byte copies when aligned) | +| `timestep_indices` | `(S,)`, `S >= 1` | int64 or int32 | In `[0, T)` | +| `token_tags` | `(S,)` | same as `timestep_indices` | In `{0, 1, 2}` | +| outputs | six `(S, H)` | `rows.dtype` | Contiguous; slices of one `(6, S, H)` buffer | + +These inputs fail closed: out-of-range tags (a fourth modality, probe H2) or timesteps +(an offset past `T`, probe H3), length or dtype mismatches, empty sequences, and tables +that are not 3 rows per timestep or 6 blocks wide. The range check is one host +read-back; pass `check_range=False` when the indices are already validated. + +## Numerics + +- **Forward** is a byte copy: one launch, one block per `(position, chunk)`, using + 16-byte vector copies. It is bitwise equal to diffusers' six `index_select` calls for + every packing, length and index dtype tested (S up to 131072). +- **Backward**: `d_rows[r] = sum over {s : r[s] = r} of grad[s]` is the op's only + reduction. Positions are sorted stably by `(row, position)` and cut into tiles of 256. + Each tile is an ascending FP32 chain from 0, and a row's tiles are left-folded in order + and rounded once. There are no atomics, so the result depends only on the indices and + the gradient. `index_select`'s own backward is an atomic scatter-add in the gradient + dtype (BF16), which is neither deterministic nor accurate. +- **Tolerance class.** The VJP is a segmented reduction, so gtest judges the op as + `reduction`. The forward is asserted bitwise in `tests/models/minimax_h3`. + +Measured on a B200 (torch 2.13.0+cu130): + +| Comparison (T = 3, S = 4097, H = 5376, BF16) | CUDA | provider (`index_select`) | +| --- | --- | --- | +| Forward vs `index_select` | bitwise equal | | +| `d_rows` repeat-bitwise | yes | no | +| `d_rows` correctly rounded from the FP64 sum | 99.99965% | 7.8% | +| `d_rows` max abs error vs FP64 | 0.25 (1 ULP) | 4.99 | + +### Fused modulation: `H3AdaLNModulationCudaOp` + +`rl_engine.backends.cuda.model_specific.minimax_h3.adaln_modulation.H3AdaLNModulationCudaOp(temb, weight, bias, +timestep_indices, token_tags)` runs [`adaln_projection_3mod`](h3-adaln-projection.md) and +this gather as one autograd node. Its forward uses the same kernels and is bitwise equal +to calling the two ops in sequence. + +The difference is in the backward. When the ops are called separately, autograd returns +the gather's FP32 segment sum to the projection in the table's BF16 dtype, which rounds it +once more. The fused node feeds the FP32 table gradient straight into the projection +backward. Use it wherever both ops run back to back, as they do in an H3 block. + +### Whole-chain backward + +`tools/validation/models/h3_chain_replay.py --backward` runs on the pinned weights for +T in {1, 2, 3, 4} × S in {3, 257, 4097, 32768}. It checks six parameter gradients in each +of the 16 cases (the time embedder's `linear_{1,2}.{weight,bias}` and block 0's AdaLN +weight and bias) against an FP64 golden in which every declared cast is straight-through: + +| S >= 257 (72 gradients) | RL-Kernel, separate ops | RL-Kernel, fused | diffusers | +| --- | --- | --- | --- | +| repeat-bitwise | 72 / 72 | 72 / 72 | 0 / 72 | +| time embedder: max abs error / golden max | 5.3e-4 to 2.4e-3 | 6.5e-7 to 4.9e-6 | 3.8e-3 to 1.3e-1, varying run to run | +| AdaLN weight/bias: correctly rounded BF16, worst case | 59.8% | 99.57% | 1.6% | + +At S = 3 all three chains are deterministic. Both RL-Kernel variants reach 1.5e-6 on the +time embedder, against 3.2e-3 for diffusers. + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op adaln_row_gather +``` + +B200, T = 3, H = 5376, BF16. "write BW" counts the 6·S·H outputs; the table stays in L2. + +| S | CUDA forward | provider forward | CUDA backward | provider backward | +| --- | --- | --- | --- | --- | +| 4097 | 0.07 ms | 0.11 ms | 0.88 ms | 1.28 ms | +| 32768 | 0.33 ms | 0.62 ms | 1.75 ms | 8.04 ms | +| 131072 | 1.17 ms | 2.42 ms | 5.62 ms | 32.45 ms | + +Backward timings and peak memory exclude leaf creation and the forward pass. The CUDA +backward stacks the six output gradients and materialises per-tile FP32 partials, adding +8130 MiB at S = 131072; the provider adds under 2 MiB. Candidate/provider execution order +alternates each iteration and is recorded in the report. + +## Evidence + +![adaln_row_gather on B200: latency and whole-chain gradient accuracy](../../reports/experiments/h3-adaln-row-gather-b200/figure.png) + +There are two data files: + +- [`report.json`](../../reports/experiments/h3-adaln-row-gather-b200/report.json): op timings, forward + bitwise checks, op-level backward, and the chain-backward cases plotted above. +- [`chain_replay.json`](../../reports/experiments/h3-adaln-row-gather-b200/chain_replay.json): the full + stage-wise forward replay and backward replay over T in {1, 2, 3, 4} × S in {3, 257, 4097, 32768}. + +`report.json` was regenerated from a clean tree at commit `80e4609` with backward-only +timings and alternating execution order. `chain_replay.json` was written from a clean tree +at commit `fa551c2`. + +## Existing implementations (RFC #420 reuse rule) + +![adaln_row_gather vs existing implementations](../../reports/experiments/h3-prior-art-b200/adaln_row_gather.png) + +| Implementation | Batch-invariant | size 4097: fwd err / worst grad err / fwd+bwd | size 32768: fwd err / worst grad err / fwd+bwd | +|---|---|---|---| +| diffusers six index_select calls | **no** (param/table grads not repeatable) | 0.0e+00 / 6.0e-02 / 1381 µs | 0.0e+00 / 2.1e-01 / 9753 µs | +| rl-kernel H3AdaLNRowGatherCudaOp | yes | 0.0e+00 / 2.5e-03 / 1106 µs | 0.0e+00 / 3.0e-03 / 2776 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +131072-row batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`adaln_row_gather.json`](../../reports/experiments/h3-prior-art-b200/adaln_row_gather.json) +was written from a clean tree at `6c900ae` by + +```bash +python tools/validation/models/h3_prior_art.py --op adaln_row_gather --out reports/experiments/h3-prior-art-b200/adaln_row_gather.json +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/adaln_row_gather.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_adaln_row_gather.py tests/models/minimax_h3/test_h3_adaln_modulation.py -v +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # whole chain, forward + backward +python tools/validation/operators/check_operator.py --op adaln_row_gather --candidate cuda --device cuda \ + --dtype bf16 --batch 3 --seq 1365 --normalized-dim 5376 --check-grad +python tools/validation/models/h3_evidence.py --op adaln_row_gather \ + --out reports/experiments/h3-adaln-row-gather-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-adaln-row-gather-b200/report.json +python tools/validation/models/h3_chain_replay.py --timesteps 1,2,3,4 --seq-lens 3,257,4097,32768 \ + --backward --out reports/experiments/h3-adaln-row-gather-b200/chain_replay.json +``` + +## Known Limitations + +- Like diffusers, the op materialises six `(S, H)` tensors, which is 8 GB at + S = 131072. Fusing the gather into the norm/modulate and gate-residual consumers + (`h3_rmsnorm`, `adaln_gate_residual`) avoids this. Those are separate rows. +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. diff --git a/docs/operators/h3-final-adaln-out.md b/docs/operators/h3-final-adaln-out.md new file mode 100644 index 000000000..f5660d315 --- /dev/null +++ b/docs/operators/h3-final-adaln-out.md @@ -0,0 +1,153 @@ +# MiniMax-H3 Final AdaLN Output + +## Summary + +`final_adaln_out` is H3's `norm_out` (`MiniMaxH3AdaLayerNormOut`), which runs once after the 50 +blocks (RFC #420, WS1 step 7): + +```text +shift, scale = norm_out.linear(silu(temb).to(bf16)).chunk(2) # (T, 5376) each, shift first +out = norm_out.norm(x) * (1.0 + scale[timestep_indices]) + shift[timestep_indices] +``` + +The table has one row per distinct timestep and is indexed by `timestep_indices`, not by +`adaln_indices`. Diffusers then upcasts `out` to the FP32 output heads' dtype. That cast is exact +and is left to the caller. + +## Entry Point + +```python +op = kernel_registry.get_op("final_adaln_out", device="cuda") +out = op(x, norm_out_norm_w, temb, norm_out_linear_w, norm_out_linear_b, timestep_indices) +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA (SM80+) | `rl_engine.backends.cuda.model_specific.minimax_h3.final_adaln_out.H3FinalAdaLNOutCudaOp` | `rl_engine._C.h3_det_linear_*`, `rl_engine._C.h3_rmsnorm_*` | One autograd node | +| PyTorch reference | `rl_engine.reference.minimax_h3.final_adaln_out.NativeH3FinalAdaLNOutOp` | n/a | `forward`: diffusers replay; `forward_fp32`: FP64 golden | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Numerics + +- **Projection.** The [AdaLN projection](h3-adaln-projection.md)'s deterministic tensor-core + GEMV (`h3-det-linear-bf16-mma-v1`), applied to `norm_out.linear`. The FP32 SiLU is rounded + once at the declared cast. BF16 `temb` is rejected (probe H7). +- **Norm and modulation.** The [`h3_rmsnorm`](h3-rmsnorm.md) kernel indexed by + `timestep_indices`. Given the same table, it is bitwise equal to diffusers. Overall, about + 0.04% of output elements differ from diffusers by 1 ULP, all traced to the projection's + summation tree versus cuBLAS. Rows are batch- and position-invariant. +- **Backward.** One autograd node, so the table gradient (FP32 segment sums) goes straight into + the projection backward without an extra BF16 rounding: + - `d_temb`, `dW` and `db` are 21–37× closer to FP64 in the recorded run than diffusers, + whose error varies from run to run because it rounds the table gradient to BF16 and accumulates it with + `index_select` atomics; + - `dx` and `d_norm_w` are at the BF16 rounding level for both; + - every gradient is repeat-bitwise. +- **Golden.** FP64, rounding only where the model stores a value in its own dtype: the SiLU + cast, `norm_out.linear`'s BF16 output table, `norm_out.norm`'s BF16 output and the BF16 + `1 + scale`. All four are rounded straight-through for the gradient. The table and `1 + scale` + are shared by every position of a timestep. Without them the golden's `d_norm_w` drifts + systematically with S, and the gtest missed it from S = 257. `norm_out.norm`'s rounding + enters `d_scale`, `dW` and `d_temb` summed over S, and the gtest missed those from S = 1024. + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op final_adaln_out +``` + +B200, B = 1, T = 3, pinned `norm_out` weights. Backward timings exclude input +and forward-graph setup: + +| S | CUDA fwd | diffusers fwd | CUDA fwd+bwd | diffusers fwd+bwd | +| --- | --- | --- | --- | --- | +| 4097 | 0.19 ms | 0.18 ms | 1.40 ms | 0.85 ms | +| 32768 | 0.46 ms | 0.80 ms | 2.59 ms | 4.49 ms | +| 131072 | 1.32 ms | 3.02 ms | 7.05 ms | 17.69 ms | + +At small S the backward is dominated by fixed setup: the stable sort for the segment sums and the +projection backward. + +The stored timing evidence includes forward graph setup in the columns labeled +`fwd+bwd`; those values are not backward-only measurements. The current benchmark +prepares the graph outside the timed backward region. Historical timings are retained +with their original commit and corrected scope. + +## Evidence + +![final_adaln_out on B200: latency and backward accuracy](../../reports/experiments/h3-final-adaln-out-b200/figure.png) + +There are two data files: + +- [`report.json`](../../reports/experiments/h3-final-adaln-out-b200/report.json): op timings, the + forward-equality fraction, row invariance and backward accuracy. Regenerated from a clean + tree at `bde1d8d` on an otherwise idle B200, with FP64 leaves and upstream gradients. The + backward reference keeps only the SiLU and table roundings, so the errors include the BF16 + rounding of `norm(x)` and `1 + scale`. +- [`chain_replay.json`](../../reports/experiments/h3-final-adaln-out-b200/chain_replay.json): the + whole conditioning chain (timestep → … → norm_out) replayed stage by stage over + T in {1, 2, 3, 4} × S in {3, 257, 4097, 32768}, with the backward replay. Written from a clean + tree at `001684d`. + +## Existing implementations (RFC #420 reuse rule) + +![final_adaln_out vs existing implementations](../../reports/experiments/h3-prior-art-b200/final_adaln_out.png) + +| Implementation | Batch-invariant | size 4097: fwd err / worst grad err / fwd+bwd | size 32768: fwd err / worst grad err / fwd+bwd | +|---|---|---|---| +| diffusers MiniMaxH3AdaLayerNormOut (op-for-op replay) | **no** (param/table grads not repeatable) | 8.3e-03 / 7.3e-02 / 1021 µs | 8.9e-03 / 2.4e-01 / 4525 µs | +| rl-kernel H3FinalAdaLNOutCudaOp | yes | 8.3e-03 / 5.2e-03 / 1642 µs | 8.9e-03 / 5.7e-03 / 2959 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +131072-token batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`final_adaln_out.json`](../../reports/experiments/h3-prior-art-b200/final_adaln_out.json) +was written from a clean tree at `03da729` by + +```bash +python tools/validation/models/h3_prior_art.py --op final_adaln_out --out reports/experiments/h3-prior-art-b200/final_adaln_out.json +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/final_adaln_out.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_final_adaln_out.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end: timestep -> ... -> norm_out +python tools/validation/operators/check_operator.py --op final_adaln_out --candidate cuda --device cuda \ + --dtype bf16 --batch 3 --seq 4097 --normalized-dim 5376 --check-grad # also 257, 1024 +python tools/validation/models/h3_evidence.py --op final_adaln_out \ + --out reports/experiments/h3-final-adaln-out-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-final-adaln-out-b200/report.json +``` + +## Known Limitations + +- **FP32 gtest at larger S.** H3 runs this op with BF16 activations and weights, and the BF16 + gtest passes at S = 257, 1024 and 4097. With `--dtype fp32` (every input FP32), the FP32 + reduction tolerance (atol = rtol = 1e-4) misses `d_temb` from S = 257 and `dW` from + S = 1024. Both gradients are FP32 sums over 3·S rows. Where large terms cancel, the + summation error exceeds the 1e-4 absolute floor. The PyTorch reference misses the same + gradients at the same S (`--batch 3`, seed 123, max abs error): + + | S | `d_temb` CUDA | `d_temb` PyTorch | `dW` CUDA | `dW` PyTorch | + | --- | --- | --- | --- | --- | + | 64 | 9.2e-5 | 1.5e-4 | 1.5e-4 | 9.2e-5 | + | 257 | 3.1e-4 ✗ | 5.9e-4 ✗ | 8.5e-4 | 2.4e-4 | + | 1024 | 1.5e-3 ✗ | 2.4e-3 ✗ | 2.0e-3 ✗ | 1.5e-3 ✗ | + | 4097 | 5.9e-3 ✗ | 9.3e-3 ✗ | 5.9e-3 ✗ | 3.9e-3 ✗ | + + ✗ fails the gtest. Unmarked values above 1e-4 pass through the rtol term. In a PyTorch + emulation, accumulating only the projection backward in FP64 clears S = 257 but not + S ≥ 1024, because the segment sums contribute as well. Making every reduction FP64 would + change the `h3_det_linear` kernels shared with `timestep_mlp_fp32` and + `adaln_projection_3mod`. The alternative is an FP32 gradient + tolerance that scales with reduction length. Which of the two to take is an open question + for the maintainers. +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. diff --git a/docs/operators/h3-one-block.md b/docs/operators/h3-one-block.md new file mode 100644 index 000000000..b95b06365 --- /dev/null +++ b/docs/operators/h3-one-block.md @@ -0,0 +1,178 @@ +# MiniMax-H3 One Block, Forward and Backward + +## Summary + +`ws1_one_h3_block` (RFC #420, WS1 closeout) runs **block 0 of the pinned MiniMax-H3 checkpoint** +node by node, forward and backward, and reports where it first departs from diffusers. The block +is a fixed graph of 17 nodes, from the AdaLN projection to the second gated residual. It runs on +packed FL2VA layouts that are **bitwise equal to the pinned diffusers pipeline's** `position_ids` +and modality tags. + +On one B200, at S = 264, 925 and 3160: + +- **Determinism.** Every node, every gradient reaching a node, and all 14 leaf gradients are + bitwise repeatable. +- **Batch invariance.** Batch row 0 of a two-row batch has the single-row bytes at every node, + forward and backward. +- **Agreement with diffusers.** Every node that promises it is bitwise equal to diffusers on + diffusers' inputs. The first drift from diffusers is `adaln_projection`, the first reduction + whose tree differs from cuBLAS. +- **Accuracy.** The block output is exactly as close to the FP32 golden as diffusers' (relative + error 5.2e-3 / 7.9e-3 / 9.0e-3). + +Only three of the 17 nodes are bound to this stack's row operators. The others use **interim** +deterministic operators until their own rows land, or the provider replay where no kernel exists +yet (see Bindings). The block-level promises above hold with whatever a node is bound to. Speed +and the exact rounding of interim nodes belong to those operators, not to this row. + +## Entry Point + +```python +from rl_engine.validation.models.h3_block import CandidateOps, run_block_backward, run_block_case +from rl_engine.validation.models.h3_weights import load_h3_block_weights + +weights = load_h3_block_weights("cuda") # conditioning + block-0 tensors, sha256-checked +ops = CandidateOps() +forward = run_block_case(weights, layout="small", ops=ops) # per node, first_*_drift +backward = run_block_backward(weights, layout="small", ops=ops) +``` + +The CLI writes the JSON evidence and the plot script renders it: + +```bash +python tools/weights/prepare_h3_weights.py --out --block # shard 1 of 14, ~4.8 GB +export RL_KERNEL_H3_WEIGHTS= +python tools/validation/models/h3_block_replay.py --layouts tiny,small,medium --out report.json +python tools/validation/models/plot_h3_block.py report.json +``` + +## Graph and Bindings + +`rl_engine/validation/models/h3_block.py` `NODES` (graph order): + +| Node | Owning RFC row | Binding today | +|---|---|---| +| `adaln_projection` | `adaln_projection_3mod` | **row**: `H3AdaLNProjectionCudaOp` (registry) | +| `norm1`, `norm2` | `h3_rmsnorm` | **row**: `H3RMSNormCudaOp.forward_modulated` (registry) | +| `q_proj`, `k_proj`, `v_proj` | `h3_qkv_gemm` | interim: Triton FP32-accumulate GEMM, one BF16 rounding | +| `q_norm`, `k_norm` | `h3_qk_rmsnorm_d128` | interim: `H3RMSNormCudaOp.forward` on the 128-wide heads | +| `rope_q`, `rope_k` | `h3_mm_rope_3axis_partial` | reference: the provider replay (no kernel exists yet) | +| `attention` | `h3_full_attention` | interim: `DeterministicAttentionOp(causal=False)` | +| `o_proj` | `h3_attention_o_gemm` | interim: Triton FP32-accumulate GEMM | +| `residual_attn`, `residual_mlp` | `adaln_gate_residual` | **row**: `H3GateResidualCudaOp` (registry) | +| `ffn_gate_up` | `h3_ffn_gate_up_gemm` | interim: Triton FP32-accumulate GEMM | +| `swiglu` | `h3_swiglu` | interim: `SwiGLUCudaOp` (FP32 math, one rounding) | +| `ffn_down` | `h3_ffn_down_gemm` | interim: Triton FP32-accumulate GEMM | + +When a row lands, replacing its node's `candidate` callable is the whole integration. The node's +report entries then become that operator's block-level evidence. + +The interim GEMMs use `_triton_gemm_fp32`, one fixed-order FP32 K loop with no split-K, and round +to BF16 once. They do not use `TritonDetGemmOp.__call__`, whose K tree has BF16 nodes. Measured on +B200: + +| Shape | Tree path err | FP32 K loop → BF16 | cuBLAS | Tree | FP32 K loop | +|---|---|---|---|---|---| +| 264 × 5376 → 7168 | 7.0e-3 | 2.4e-3 | 2.4e-3 | 10.4 ms | 1.2 ms | +| 925 × 14336 → 5376 | 7.5e-3 | 2.6e-3 | 2.6e-3 | 47 ms | 2.7 ms | + +Both paths are bitwise repeatable and row-invariant. Error is max|err| / max|ref| against FP64. + +## What Each Node Reports + +Every node runs three ways: the provider replay (`h3_provider`, diffusers@f53d552 op for op), the +candidate (chained, and *isolated* on the provider's inputs), and the FP32 golden +(`rl_engine/reference/minimax_h3/block.py`: FP32 math, TF32 off, no BF16 storage between nodes). + +- **Forward:** `repeat_bitwise_equal` and `batch_row_bitwise_equal` (row 0 of a B = 2 batch whose + row 1 has a different input). Also chained and isolated agreement with the provider (bitwise, + max abs, first differing packed row), the candidate's and provider's `max|err| / max|golden|`, + and the sha256 of the node's output. +- **Backward:** for the block input, `temb` and the 12 block parameters: repeat bitwise, and the + candidate's and provider's error against the golden. For the gradient reaching each node's + output, in reverse graph order: repeat bitwise and batch-row bitwise. Parameter gradients sum + over the batch, so only activation gradients take part in the batch check. +- **First drift:** `first_drift`, `first_isolated_drift`, `first_repeat_drift`, + `first_batch_drift`, `first_grad_repeat_drift` and `first_grad_batch_drift` each name the first + node where that comparison stops being bitwise. +- **Permutation probe (control H-C3, reported, not promised):** permute the packed rows together + with positions and AdaLN indices, and compare the output to the permuted original. + +Inputs: `temb` comes from the provider's conditioning path (the conditioning rows are checked by +`h3_chain`). The block input is a seeded BF16 stand-in, because the packed input projections are +other rows. Text rows use AdaLN table row 0 and media rows use row 1. + +## Evidence + +[`report.json`](../usage/evidence/h3-one-block-b200/report.json) was generated from a clean tree +at `41e4293` on an otherwise idle B200 (torch 2.13.0+cu130): + +![one-block evidence](../usage/evidence/h3-one-block-b200/figure.png) + +| Layout | S | First drift / first isolated drift | Repeat / batch drift (fwd, grad) | Output err: RL-Kernel / diffusers | Fwd ms: RL-Kernel / diffusers | Fwd+bwd ms | +|---|---|---|---|---|---|---| +| tiny | 264 | `adaln_projection` / `adaln_projection` | none | 5.21e-3 / 5.21e-3 | 7.4 / 0.61 | 22.7 / 2.8 | +| small | 925 | `adaln_projection` / `adaln_projection` | none | 7.87e-3 / 7.87e-3 | 30.1 / 1.12 | 81.7 / 5.0 | +| medium | 3160 | `adaln_projection` / `adaln_projection` | none | 8.97e-3 / 8.97e-3 | 225 / 3.8 | 577 / 11.5 | + +- **Forward accuracy.** Per node, the candidate's error is within 0.93–1.20× diffusers'. The + largest is `attention` at S = 3160, which is the interim attention operator's. +- **Gradient accuracy.** The candidate is better than diffusers on `adaln_proj.weight`, + `adaln_proj.bias` and `temb`, by up to 7×: diffusers' `index_select` backward accumulates the + table gradient in BF16 atomics, while the row operators use FP32 segment sums. The candidate is + within 1.16× on the other leaves, except `attn.norm_q.weight` at S = 3160 (9.5e-3 vs 5.5e-3). + At S = 925 the same leaf is 2× *more* accurate than diffusers (1.3e-2 vs 2.4e-2), so this + 128-element gradient's error moves with BF16 rounding rather than favouring one side. +- **Permutation probe (H-C3).** The output changes in 79 / 450 / 2199 of the rows, by at most 8. + Attention sums over keys in packed order, so a consistent permutation is not byte-preserving with + the interim attention operator. This is for `h3_full_attention` to settle. +- **Speed.** The candidate is 12–60× slower than diffusers forward. Almost all of it is the + interim GEMMs and attention (FP32 scores materialised). The three row operators are each + faster than or equal to diffusers in their own evidence. + +## Existing implementations (RFC #420 reuse rule) + +- **Is there an existing block-level harness?** diffusers' `MiniMaxH3TransformerBlock` is the + executable reference. It is replayed op for op as the provider, so diffusers is not imported and + the replay pins the reference's dtypes and ordering. The Qwen3-8B chain gate from #315 + (`rl_engine/validation/models/chain_gate.py`) has per-node token digests and first-drift + localisation, but its graph, shapes and acceptance are Qwen3-specific (causal attention, no + AdaLN). This row reuses its method (per-node digests, graph-ordered first drift) on the H3 + `h3_chain` stage pattern. +- **Interim operators.** These are reused from RL-Kernel and are not new code. The only choice + made here is the FP32-accumulate entry of the Triton GEMM instead of its BF16 tree, for the + accuracy and speed measured above. + +## Tests + +`tests/models/minimax_h3/test_h3_one_block.py`: + +- **CPU:** + - the three layouts hash to the pinned diffusers `build_packed_sequence` output; + - the graph is well formed; + - the partial MM-RoPE leaves channels 96–127 unchanged. +- **GPU, `tiny` and `small` forward:** + - every node repeats and is batch-row bitwise; + - promised nodes are isolated-bitwise to diffusers; + - node and output errors are at most 1.25 × diffusers' + 1e-3. +- **GPU, `tiny` backward:** + - every leaf repeats; + - dX is batch-row bitwise; + - no gradient drift; + - leaf errors are at most 1.5 × diffusers' + 5e-3. + +```bash +RL_KERNEL_H3_WEIGHTS= CUDA_VISIBLE_DEVICES=0,1 python -m pytest tests/models/minimax_h3 -q +``` + +## Known Limitations + +- **Single platform.** CUDA only. `DeterministicAttentionOp` and the Triton GEMM also build for + ROCm, but the block has not been run there. +- **Interim nodes.** Eleven of the 17 nodes are interim or reference bindings. Their speed and + rounding are not the owning rows' final ones, and the MM-RoPE nodes are the provider replay + itself, so they carry no RL-Kernel evidence until `h3_mm_rope_3axis_partial` lands. +- **Stand-in inputs.** The block input is a seeded stand-in and `temb` is the provider's. + Real packed projections arrive with the input-projection rows. +- **Not covered here.** Block 0 only (`ws1_full_h3_transformer` covers 50 blocks and the refiner), + and one sequence per batch row (H3's batch is a replication axis). diff --git a/docs/operators/h3-rmsnorm.md b/docs/operators/h3-rmsnorm.md new file mode 100644 index 000000000..c22588cd5 --- /dev/null +++ b/docs/operators/h3-rmsnorm.md @@ -0,0 +1,150 @@ +# MiniMax-H3 RMSNorm and AdaLN Modulation + +## Summary + +`h3_rmsnorm` covers every RMSNorm in MiniMax-H3 (RFC #420, WS1 step 4). That is the block +`norm1`/`norm2`, the token-refiner norms, the refiner `final_norm` and `norm_out.norm`. All of +them are `nn.RMSNorm(5376, eps=1e-5)` with a BF16 affine weight. In a transformer block and in +`norm_out`, the normalised rows are immediately modulated by per-row AdaLN parameters, and the +op fuses that step: + +```text +n = rms_norm(x, weight, eps) +out = n * (1.0 + scale[index]) + shift[index] +``` + +`index` is `adaln_indices` in a block (rows of the [projection](h3-adaln-projection.md) table) +and `timestep_indices` in `norm_out`. `shift`/`scale` are `(R, H)` row views of the AdaLN table, +gathered inside the kernel, so the `(S, H)` tensors that +[`adaln_row_gather`](h3-adaln-row-gather.md) materialises are never needed. + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("h3_rmsnorm", device="cuda") +n = op(x, weight) # plain RMSNorm +out = op.forward_modulated(x, weight, shift_msa, scale_msa, adaln_indices) +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA (SM80+, validated on SM100) | `rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp` | `rl_engine._C.h3_rmsnorm_{forward,backward}` | Bitwise equal to `nn.RMSNorm` and to the diffusers modulation | +| PyTorch reference | `rl_engine.reference.minimax_h3.rmsnorm.NativeH3RMSNormOp` | n/a | `forward*`: provider path; `forward*_fp32`: FP64 golden (the modulated one stores `norm(x)` and `1 + scale` in BF16, straight-through, as the model does) | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `x` | `(..., S, N)` | bf16 / fp16 / fp32 | `N % 4 == 0`; H3: `N = 5376` | +| `weight` | `(N,)` | `x.dtype` | | +| `shift`, `scale` | `(R, N)` row views, same stride | `x.dtype` | Unit column stride | +| `index` | `(S,)` | int64 | In `[0, R)`; one entry per position, shared across the batch | + +Out-of-range indices, mismatched dtypes or shapes, a non-positive eps, and `N % 4 != 0` all fail +closed. + +## Numerics + +- **Statistics (contract `h3-rmsnorm-v1`).** The kernel replays PyTorch's own + `vectorized_layer_norm_kernel` (torch `cf30153`): + - one `(32, 4)` block per row, reading 4-element vectors; + - thread `t` sums vectors `t, t+128, …` in order; + - a shuffle-down tree (16 → 1), then a cross-warp tree; + - `rsqrtf(Σx²/N + eps)`, then `w * (rstd * x)` with one cast. + + The plain norm is therefore **bitwise equal to `nn.RMSNorm`**. Rows are independent of batch + size and position. +- **Modulation.** `1 + scale`, `n * (…)` and `+ shift` are each rounded to the tensor dtype, + exactly where the eager expression rounds. The fused output is **bitwise equal to diffusers' + `norm(x) * (1.0 + scale.index_select(0, i)) + shift.index_select(0, i)`**. +- **Backward.** Everything is in FP32 with one cast per output and no atomics: + - `dx` is row-local, using the same reduction tree; + - `dweight` is folded over fixed 256-row tiles in ascending order; + - `dshift`/`dscale` are segmented sums over positions sorted stably by table row, in the + same scheme as [`adaln_row_gather`](h3-adaln-row-gather.md). + + Diffusers' `index_select` backward accumulates `dshift`/`dscale` with BF16 atomics, so it is + non-deterministic and about 15–20× less accurate (its error varies from run to run). + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op h3_rmsnorm +``` + +B200, block `norm1` + MSA modulation, B = 1, H = 5376, BF16: + +| S | CUDA fwd | diffusers fwd | CUDA bwd | diffusers bwd | +| --- | --- | --- | --- | --- | +| 4097 | 0.08 ms | 0.12 ms | 1.23 ms | 0.85 ms | +| 32768 | 0.35 ms | 0.76 ms | 2.18 ms | 4.40 ms | +| 131072 | 1.21 ms | 2.90 ms | 6.17 ms | 17.55 ms | + +The forward is a single pass that never materialises the gathered rows. At small S the backward +is dominated by the fixed cost of the stable sort and tile setup. Backward timings and peak +memory exclude leaf creation and the forward pass. Candidate/provider execution order alternates +each iteration and is recorded in the report. + +## Evidence + +![h3_rmsnorm on B200: latency and backward accuracy](../../reports/experiments/h3-rmsnorm-b200/figure.png) + +The data is in [`report.json`](../../reports/experiments/h3-rmsnorm-b200/report.json), written by +`tools/validation/models/h3_evidence.py` from a clean tree at commit `80e4609`. It also records: + +- bitwise equality with `nn.RMSNorm` for all four pinned norm weights; +- bitwise equality with diffusers for the modulation; +- row invariance. + +## Existing implementations (RFC #420 reuse rule) + +![norm_modulate vs existing implementations](../../reports/experiments/h3-prior-art-b200/norm_modulate.png) + +| Implementation | Batch-invariant | size 4097: fwd err / worst grad err / fwd+bwd | size 32768: fwd err / worst grad err / fwd+bwd | +|---|---|---|---| +| diffusers composition (F.rms_norm + index_select modulation) | **no** (param/table grads not repeatable) | 8.6e-03 / 4.5e-02 / 872 µs | 7.5e-03 / 1.9e-01 / 4434 µs | +| torch F.rms_norm (no modulation) | yes | 2.1e-03 / 2.1e-03 / 331 µs | 2.0e-03 / 2.8e-03 / 850 µs | +| TE 2.20.2 RMSNorm (no modulation) | yes | 2.1e-03 / 2.1e-03 / 471 µs | 2.0e-03 / 2.8e-03 / 1597 µs | +| Liger 0.8.4 modulated RMSNorm + index_select | **no** (param/table grads not repeatable) | 6.6e-03 / 4.7e-02 / 810 µs | 6.9e-03 / 2.0e-01 / 3911 µs | +| Liger 0.8.4 modulated RMSNorm + rl-kernel row gather | yes | 6.6e-03 / 6.5e-03 / 1824 µs | 6.9e-03 / 8.0e-03 / 4763 µs | +| SGLang 0.5.21 fused_norm_scale_shift (forward only) | yes | 4.4e-03 / — / 85 µs (fwd only) | 4.1e-03 / — / 436 µs (fwd only) | +| rl-kernel H3RMSNormCudaOp.forward_modulated | yes | 8.6e-03 / 4.2e-03 / 1093 µs | 7.5e-03 / 5.0e-03 / 2653 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +131072-token batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`norm_modulate.json`](../../reports/experiments/h3-prior-art-b200/norm_modulate.json) +was written from a clean tree at `ee83dec` by + +```bash +python tools/validation/models/h3_prior_art.py --op norm_modulate --out reports/experiments/h3-prior-art-b200/norm_modulate.json +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/norm_modulate.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_rmsnorm.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end, incl. block norm1 + modulation +python tools/validation/operators/check_operator.py --op h3_rmsnorm --candidate cuda --device cuda \ + --dtype bf16 --batch 2 --seq 257 --normalized-dim 5376 --check-grad +python tools/validation/models/h3_evidence.py --op h3_rmsnorm --out reports/experiments/h3-rmsnorm-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-rmsnorm-b200/report.json +``` + +## Known Limitations + +- Bitwise parity with `nn.RMSNorm` is tied to PyTorch's vectorized layer-norm path (`N % 4 == 0`, + aligned tensors). Other shapes are rejected rather than approximated. +- The token-refiner blocks' own norm weights live in shard 14, which is not pinned. They are the + same module (`nn.RMSNorm(5376, eps=1e-5)`) as the four pinned norms tested here. +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. diff --git a/docs/operators/h3-sp-norm-adaln.md b/docs/operators/h3-sp-norm-adaln.md new file mode 100644 index 000000000..1f19d0444 --- /dev/null +++ b/docs/operators/h3-sp-norm-adaln.md @@ -0,0 +1,117 @@ +# MiniMax-H3 Sequence-Parallel Norm, AdaLN Modulation and Gated Residual + +## Summary + +`sp_norm_adaln` runs the row-wise part of an H3 block with the packed sequence split across +sequence-parallel ranks (RFC #420, WS2). It covers the [RMSNorm and AdaLN modulation](h3-rmsnorm.md) +and the [gated residual](h3-adaln-gate-residual.md). Every rank's output and row gradients are the +**WS1 bytes of its rows**. The gradients of the replicated tensors (`d_norm_weight`, `d_shift`, +`d_scale`, `d_gate`) are **the WS1 bytes, identical on every rank**. + +## Entry Point + +```python +from rl_engine.distributed.algorithms.collectives import DeterministicCollective +from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import H3SPNormAdaLNCudaOp + +collective = DeterministicCollective(group=sp_group, device=local_rank) +op = H3SPNormAdaLNCudaOp(collective, seq_len=S, batch=B) +rows = slice(op.layout.lo, op.layout.hi) # this rank's packed positions +hidden = op.gate_residual(residual[:, rows], attn_out_local, gate_msa, adaln_indices) +normed = op.norm_modulated(hidden, norm2_weight, shift_mlp, scale_mlp, adaln_indices) +final = op.norm_modulated(h, norm_out_weight, shift, scale, timestep_indices) # norm_out +op.readback() # sp, rank, positions, collective backend, reduction order, fallback +``` + +The table views, the norm weight and the full `(S,)` row index are replicated. Activations are +the rank's `(B, S_local, H)` rows. + +## Ownership + +Rank `r` of `sp` holds packed positions `[r * S // sp, (r + 1) * S // sp)` of every batch item, +together with the matching rows of the metadata. Any `S >= sp` works, including `S` not divisible +by `sp`. Shapes that do not match the rank's slice, a sharded index, an out-of-range index, CPU +tensors and malformed tables all fail closed. + +## Numerics + +- **Forward and row gradients.** `out`, `dx`, `d_sublayer` and `d_residual` are the WS1 kernels on + the local rows, with the local slice of the index. Rows are independent, so they are byte-equal + to WS1. +- **Cross-row gradients.** In WS1, `d_norm_weight` sums `d_n * x * rstd` over fixed 256-row tiles + of the `B x S` rows. `d_shift`/`d_scale`/`d_gate` sum over 256-element tiles of each table row's + positions, sorted stably. In both cases the tile partials are then folded in ascending order. SP + keeps exactly that two-level order: + 1. Every rank builds the global WS1 tile list from the replicated index (`sp_plan`). + 2. Each tile is computed by the rank that holds its first row. Rows of that tile that live on + other ranks are all-gathered first. Only tiles that cross a shard boundary move rows. + 3. The tile partials (`h3_rmsnorm_backward_partials`, `h3_gate_grad_partials`, the WS1 partial + kernels over explicit row lists) are all-gathered and put back in WS1 tile order. + 4. Every rank runs the WS1 folds (`h3_rmsnorm_fold_partials`, `h3_gate_grad_fold`). + + The only collectives are rank-ordered all-gathers, which are copies. No collective reduction + arithmetic is used. + +A naive SP backward runs the WS1 backward on each rank's rows and sums the per-rank results. It +changes the reduction tree, so it is not WS1. The evidence measures how often it differs. + +The plan depends only on the layout and the index. `H3SPNormAdaLNCudaOp` builds it once per index +object and reuses it for every norm and gated residual that uses the same index, which in a forward +pass is every block. + +## Rows Exchanged + +How many rows move depends on how the table rows interleave along the sequence: + +| Packing | Rows sent per rank (S = 32768, SP8, 4096 rows per rank) | +| --- | --- | +| block (H3's layout: each timestep's text, video and audio tokens contiguous) | 0 – 224 | +| interleaved (modality and timestep random per position; stress case) | 0 – 1536 | + +## Evidence + +![sp_norm_adaln on 8 x B200: byte equality, time, naive SP](../../reports/experiments/h3-sp-norm-adaln-b200/figure.png) + +The report is written by `tools/validation/models/h3_ws2_evidence.py` from a clean tree at commit `6c0d380`, with +real NCCL processes on one 8 x B200 node, one GPU per rank. The region is +`norm2(residual + gate_msa[row] * y)` with `shift_mlp`/`scale_mlp` modulation, H = 5376, BF16. It +covers six cases: S = 4097 (block, interleaved, B = 2), 32768 (block, interleaved) and 131072 +(block). For SP 1, 2, 4 and 8, every rank's output rows, `d_residual` and `d_sublayer`, and its +`d_norm_w` and `d_table`, are byte-equal to WS1 computed on the same GPU. + +Forward + backward of the region (slowest rank): + +| S, packing | WS1 (1 GPU) | SP2 | SP4 | SP8 | +| --- | --- | --- | --- | --- | +| 4097, block | 2.04 ms | 1.73 ms | 1.88 ms | 2.43 ms | +| 32768, block | 3.72 ms | 2.85 ms | 2.40 ms | 2.72 ms | +| 32768, interleaved | 3.72 ms | 3.45 ms | 3.88 ms | 6.82 ms | +| 131072, block | 10.73 ms | 7.27 ms | 4.66 ms | 3.92 ms | + +At long sequences the rows split and the gain grows: 2.74x at SP8 for S = 131072. At short +sequences, each backward's handful of host-synchronising all-gathers (about 0.1 ms each) +dominates. Interleaved packing also moves more rows. A naive SP8 backward (each rank's WS1 +backward, then a rank-order sum) differs from WS1 on 40% of `d_norm_w` elements and on 6% (block) +or 21% (interleaved) of `d_table` elements. + +## Tests + +```bash +python -m pytest tests/models/minimax_h3/test_h3_sp_norm_adaln.py -v # layout, rank backward, NCCL SP2/4/8 +python tools/validation/models/h3_ws2_evidence.py --op sp_norm_adaln --worlds 1,2,4,8 \ + --out reports/experiments/h3-sp-norm-adaln-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-sp-norm-adaln-b200/report.json +``` + +The rank-backward tests run SP 2–8 on one GPU, with ranks as threads around the plain backward +functions. They cover odd `S`, `B = 2`, shards smaller than a tile, block and interleaved packing, +and bf16 and fp32, and compare every rank against WS1. The NCCL tests run the autograd region end +to end and need as many GPUs as ranks. + +## Known Limitations + +- CUDA only. `DeterministicCollective` supports world sizes 1, 2, 4 and 8. +- With interleaved packing, up to a full shard of rows can move in the backward. The result is + still exact. +- Attention and FFN are other rows. Until they land, the region uses a seeded stand-in for the + sublayer output. diff --git a/docs/operators/h3-timestep-mlp.md b/docs/operators/h3-timestep-mlp.md new file mode 100644 index 000000000..7f2d477e3 --- /dev/null +++ b/docs/operators/h3-timestep-mlp.md @@ -0,0 +1,165 @@ +# MiniMax-H3 FP32 Timestep MLP + +## Summary + +`timestep_mlp_fp32` is H3's `TimestepEmbedding(in_channels=256, time_embed_dim=5376, +out_dim=2688)`. It runs on the sinusoidal features of the distinct timesteps +(RFC #420, WS1 step 3): + +```text +temb = linear_2(silu(linear_1(features))) (T, 256) -> (T, 5376) -> (T, 2688) +``` + +`time_embedder` is a `_keep_in_fp32_modules` module. Its weights, biases, +activations and output are all FP32. `temb` stays FP32 because every AdaLN +projection applies its own SiLU before casting to BF16. + +Pinned model: `MiniMaxAI/MiniMax-H3@42ed227`. The tensors are +`time_embedder.linear_{1,2}.{weight,bias}` (see `rl_engine/validation/models/h3_manifest.json`). + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("timestep_mlp_fp32", device="cuda") +temb = op(features, w1, b1, w2, b2) # all float32 +``` + +## Backends + +| Backend | Wrapper | Native symbols | Status | +| --- | --- | --- | --- | +| CUDA (SM90, SM100) | `rl_engine.backends.cuda.model_specific.minimax_h3.timestep_mlp.H3TimestepMLPCudaOp` | `rl_engine._C.h3_det_linear_{forward,backward_input,backward_weight}` | Contract `h3-det-linear-v1` | +| PyTorch reference | `rl_engine.reference.minimax_h3.timestep_mlp.NativeH3TimestepMLPOp` | n/a | `forward`: provider path (`F.linear`/`F.silu`); `forward_fp32`: FP64 golden | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `x` | `(T, K)`, `T >= 1` | float32 | H3: `K = 256` | +| `w1`, `b1` | `(H, K)`, `(H,)` | float32 | H3: `H = 5376` | +| `w2`, `b2` | `(D, H)`, `(D,)` | float32 | H3: `D = 2688` | +| output | `(T, D)` | float32 | | + +The CUDA kernel needs `K` and `H` to be multiples of 4, for 16-byte vector loads. +These inputs fail closed: BF16 in any argument (RFC probe H7), mismatched shapes, +an empty `x`, mixed devices, and unaligned `K`. + +## Numerics: contract `h3-det-linear-v1` + +Forward, for each output `y[t, n]`: + +- One warp owns a group of output columns. Lane `l` reads the 16-byte K chunks + `l, l + 32, l + 64, ...`. +- Each lane accumulates its chunks in ascending order, and the elements inside a chunk + in ascending order, using `fmaf` into an FP32 accumulator that starts at 0. +- The 32 lane sums are combined with an xor butterfly (16, 8, 4, 2, 1). +- The bias is added once, then SiLU runs in FP32 as `v / (1 + expf(-v))` + (PyTorch's formula). The result is stored as FP32. + +How many columns a warp owns, the row tile (1, 2 or 4), and how many chunks are +loaded ahead change only when loads are issued, never the order in which a column +is summed. A timestep's `temb` is therefore bitwise identical however many other +timesteps share the call, and wherever it sits among them. + +Backward is deterministic. There is no cuBLAS and there are no atomics: + +- `d_hidden = g @ W2` and `dx = d_pre @ W1`: N is split into fixed 64-row chunks. + Each chunk is an ascending `fmaf` chain, and the chunks are then left-folded in + ascending order. These gradients are row-local and batch-invariant. +- `d_pre = d_hidden * s * (1 + z * (1 - s))`, elementwise in FP32, using the saved + pre-activation `z`. +- `dW[n, k] = sum_t g[t, n] * x[t, k]` and `db[n] = sum_t g[t, n]`: an ascending-`t` + FP32 fold that starts from 0, in logical row order. + +Accuracy, measured on a B200 with the pinned checkpoint weights, TF32 disabled and 200 +draws of 4 timesteps (each draw includes t = 0 and t = 1): + +| Comparison | CUDA | provider (cuBLAS) | +| --- | --- | --- | +| max abs error vs FP64 golden, median over draws | 1.1e-6 | 5.4e-7 | +| gradients vs FP64 autograd (`x`, `w1`, `b1`, `w2`, `b2`) | <= 1.2e-5 | | + +Both outputs are about 100× inside the contract (`reduction` / `float32`, atol and rtol +1e-4). Neither is more accurate in general: cuBLAS's tree happens to do slightly better +on these weights. What the CUDA kernel adds is a fixed summation order, which makes each +row batch- and position-invariant and every run repeat-bitwise, and it is 2× faster +for T >= 2. The two outputs are not bitwise equal, so `tools/validation/models/h3_chain_replay.py` +reports this stage as the chain's `first_drift`, which is expected for a reduction. + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op timestep_mlp_fp32 +``` + +B200, pinned weights (63.3 MB FP32). End-to-end op times include the Python wrapper: + +| T | CUDA op | provider | kernels only (CUDA) | +| --- | --- | --- | --- | +| 1 | 35.5 µs | 40.8 µs | 3.8 + 11.6 µs | +| 2 | 40.4 µs | 84.9 µs | 3.9 + 13.7 µs | +| 4 | 44.9 µs | 88.6 µs | 5.0 + 15.8 µs | + +The 5376→2688 layer reads its weights at about 5.0 TB/s at T = 1. For T >= 2 the +provider switches from cuBLAS GEMV to an SGEMM path, which takes about 52 µs of +kernel time. + +## Evidence + +![timestep_mlp_fp32 on B200: latency and per-draw error vs FP64](../../reports/experiments/h3-timestep-mlp-b200/figure.png) + +The data is in [`report.json`](../../reports/experiments/h3-timestep-mlp-b200/report.json), written +by `tools/validation/models/h3_evidence.py` from a clean tree at commit `65ef7f6`. It also records that a +timestep's row is bitwise identical whether it runs alone or in a batch of 9. + +## Existing implementations (RFC #420 reuse rule) + +![timestep_mlp vs existing implementations](../../reports/experiments/h3-prior-art-b200/timestep_mlp.png) + +| Implementation | Batch-invariant | size 3: fwd err / worst grad err / fwd+bwd | size 256: fwd err / worst grad err / fwd+bwd | size 2048: fwd err / worst grad err / fwd+bwd | +|---|---|---|---|---| +| diffusers TimestepEmbedding, FP32 F.linear [plain] | **no** (14214 rows, 4108 sub-batches, 2232 sweep cases) | 5.7e-07 / 4.1e-07 / 777 µs | 5.7e-07 / 1.0e-06 / 832 µs | 3.7e-06 / 2.0e-06 / 3609 µs | +| rl-kernel H3TimestepMLPCudaOp | yes | 2.6e-07 / 4.2e-07 / 748 µs | 2.3e-07 / 7.6e-07 / 4595 µs | 2.6e-07 / 2.3e-06 / 42993 µs | +| diffusers TimestepEmbedding, FP32 F.linear [vllm] | **no** (14214 rows, 4096 sub-batches, 2232 sweep cases) | 1.8e-07 / 4.5e-07 / 752 µs | 3.7e-06 / 2.8e-06 / 921 µs | 3.7e-06 / 3.2e-06 / 3826 µs | +| diffusers TimestepEmbedding, FP32 F.linear [sglang] | yes | 1.6e-03 / 1.6e-03 / 1036 µs | 1.5e-03 / 1.5e-03 / 1140 µs | 1.5e-03 / 1.7e-03 / 2636 µs | +| diffusers TimestepEmbedding, FP32 F.linear [sglang_ieee] | yes | 3.3e-06 / 2.3e-06 / 2234 µs | 3.7e-06 / 2.8e-06 / 2428 µs | 3.7e-06 / 3.2e-06 / 7329 µs | +| diffusers TimestepEmbedding, FP32 F.linear [megatron_te_native] | **no** (14214 rows, 4108 sub-batches, 2232 sweep cases) | 1.8e-07 / 4.1e-07 / 729 µs | 3.7e-06 / 1.0e-06 / 814 µs | 3.7e-06 / 2.0e-06 / 3636 µs | +| diffusers TimestepEmbedding, FP32 F.linear [megatron_triton] | yes | 1.6e-03 / 1.6e-03 / 1027 µs | 1.5e-03 / 1.5e-03 / 1137 µs | 1.5e-03 / 1.7e-03 / 2634 µs | +| diffusers TimestepEmbedding, FP32 F.linear [megatron_triton_ieee] | yes | 3.3e-06 / 2.3e-06 / 2309 µs | 3.7e-06 / 2.8e-06 / 2435 µs | 3.7e-06 / 3.2e-06 / 7349 µs | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +4096-timestep batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`timestep_mlp.json`](../../reports/experiments/h3-prior-art-b200/timestep_mlp.json) +was written from a clean tree at `ffd958a` by + +```bash +python tools/validation/models/h3_prior_art.py --op timestep_mlp --out reports/experiments/h3-prior-art-b200/timestep_mlp.json --megatron-src +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/timestep_mlp.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_timestep_mlp.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end: sinusoid -> MLP +python tools/validation/operators/check_operator.py --op timestep_mlp_fp32 --candidate cuda --device cuda \ + --dtype fp32 --batch 3 --check-grad +python tools/validation/models/h3_evidence.py --op timestep_mlp_fp32 \ + --out reports/experiments/h3-timestep-mlp-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-timestep-mlp-b200/report.json +``` + +## Known Limitations + +- FP32 only, by contract. The BF16 and FP16 gtest rows do not apply to this op. +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. +- Weight gradients are reductions across timesteps, so they follow logical row + order. They are deterministic, but not batch-invariant (no cross-row reduction can be). diff --git a/docs/operators/h3-timestep-sinusoid.md b/docs/operators/h3-timestep-sinusoid.md new file mode 100644 index 000000000..7bf0f13b7 --- /dev/null +++ b/docs/operators/h3-timestep-sinusoid.md @@ -0,0 +1,137 @@ +# MiniMax-H3 Timestep Sinusoid + +## Summary + +`timestep_sinusoid_h3` produces the FP32 sinusoidal timestep features that feed the +MiniMax-H3 timestep MLP (RFC #420, WS1 step 3). H3 builds +`Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)` and calls it +on the *distinct* timesteps of the packed sequence, unscaled in `[0, 1]` +(`t = 1 - sigma`): + +```text +freq[k] = exp(-ln(10000) * k / 128) k = 0..127 +arg[t, k] = t * freq[k] +out[t] = [cos(arg[t]) | sin(arg[t])] (T, 256), float32 +``` + +Pinned model: `MiniMaxAI/MiniMax-H3@42ed227` (`rl_engine/validation/models/h3_manifest.json`). +Provider reference: `huggingface/diffusers@f53d552`, `get_timestep_embedding`. + +## Entry Point + +```python +from rl_engine.runtime.registry import kernel_registry + +op = kernel_registry.get_op("timestep_sinusoid_h3", device="cuda") +features = op(timestep) # timestep: (T,) in [0, 1] -> (T, 256) float32 +``` + +## Backends + +| Backend | Wrapper | Native symbol | Status | +| --- | --- | --- | --- | +| CUDA (SM90, SM100) | `rl_engine.backends.cuda.model_specific.minimax_h3.timestep_sinusoid.H3TimestepSinusoidCudaOp` | `rl_engine._C.h3_timestep_sinusoid_forward` | Bitwise equal to the provider path | +| PyTorch reference | `rl_engine.reference.minimax_h3.timestep_sinusoid.NativeH3TimestepSinusoidOp` | n/a | `forward`: provider replay; `forward_fp32`: FP64 golden | +| ROCm | n/a | n/a | Falls back to the PyTorch reference | + +## Tensor Contract + +| Argument | Shape | Dtype | Requirements | +| --- | --- | --- | --- | +| `timestep` | `(T,)`, `T >= 1` | fp32 / bf16 / fp16 | Finite, in `[0, 1]`; any stride; CUDA for the CUDA backend | +| `num_channels` | scalar | int | Positive and even; H3 uses 256 | +| output | `(T, num_channels)` | float32 | `[cos \| sin]` order | + +Inputs that break the contract fail closed: non-1-D, empty, integer or out-of-range +timesteps raise. `t * 1000` callers (RFC probe H10) therefore fail instead of +silently producing a different embedding. The range check costs one host read-back. +The native CUDA entrypoint performs the value check, so the wrapper does not repeat +the host synchronization. Pass `check_range=False` to `forward` or +`rl_engine._C.h3_timestep_sinusoid_forward` only when the caller has already validated +`t`; this trusted path skips value validation and is used for kernel-only profiling. + +## Numerics + +- **Operation order.** The kernel evaluates exactly the FP32 operations the provider + runs on CUDA: `(float)(-ln 10000) * k`, multiplied by the FP32 reciprocal of `128` + (PyTorch divides by a CPU scalar this way), `expf`, `t * freq`, `cosf` and `sinf`. + The multiplies use `__fmul_rn`, so they are never fused into an FMA. The extension + is built without `--use_fast_math`, so `expf`, `sinf` and `cosf` are the precise + libdevice functions. +- **Cast points.** bf16 and fp16 timesteps are upcast to FP32 before anything else, + matching `timesteps[:, None].float()`. The output is never rounded below FP32. +- **Invariance.** Each output element depends only on `(t, k)`. Results are + bitwise identical across batch size, position, permutation and repeated runs. +- **Backward.** `d/dt` is the analytic row-local VJP + (`-f * sin(t f)` for the cos half, `f * cos(t f)` for the sin half). It is summed over + the 256 channels with a fixed pairwise tree (`tree_sum_lastdim_fp32`), so it does not + depend on how many timesteps share the launch. + +Measured on a B200 (torch 2.13.0+cu130): + +| Comparison | Result | +| --- | --- | +| CUDA vs provider path, T in {1, 2, 3, 7, 64, 1000, 4097} | bitwise equal | +| CUDA vs FP64 golden | max abs 1.19e-7 (contract: FP32 elementwise atol 1e-5) | +| gradient vs FP64 golden (gtest, T = 7) | max abs 2.4e-7 | + +## Performance Notes + +```bash +python benchmarks/models/benchmark_h3_conditioning.py --op timestep_sinusoid_h3 +``` + +On a B200 this is a single launch of about 15 µs, independent of `T` for `T <= 64`. With +the range check it takes about 49 µs. The provider path takes about 62–70 µs, because it is +six eager kernels plus two concatenations. Either way the op is launch-bound: it moves +about 1 KB per timestep. + +## Evidence + +![timestep_sinusoid_h3 on B200: latency and error vs FP64](../../reports/experiments/h3-timestep-sinusoid-b200/figure.png) + +The data is in [`report.json`](../../reports/experiments/h3-timestep-sinusoid-b200/report.json). +`tools/validation/models/h3_evidence.py` wrote it from a clean tree at commit `0522865`, and the report +records that commit and the environment. + +## Existing implementations (RFC #420 reuse rule) + +![timestep_sinusoid vs existing implementations](../../reports/experiments/h3-prior-art-b200/timestep_sinusoid.png) + +| Implementation | Batch-invariant | size 3: fwd err / worst grad err / fwd+bwd | size 256: fwd err / worst grad err / fwd+bwd | size 2048: fwd err / worst grad err / fwd+bwd | +|---|---|---|---|---| +| diffusers get_timestep_embedding (op-for-op replay) | yes | 6.7e-08 / — / 64 µs (fwd only) | 8.6e-08 / — / 64 µs (fwd only) | 9.9e-08 / — / 65 µs (fwd only) | +| SGLang 0.5.21 timestep_embedding | yes | 6.7e-08 / — / 13 µs (fwd only) | 8.6e-08 / — / 12 µs (fwd only) | 9.9e-08 / — / 13 µs (fwd only) | +| rl-kernel H3TimestepSinusoidCudaOp (check_range=False) | yes | 6.7e-08 / — / 16 µs (fwd only) | 8.6e-08 / — / 16 µs (fwd only) | 9.9e-08 / — / 16 µs (fwd only) | + +Errors are max|err| / max|ref| against the same computation in FP64; latency is the median +forward + backward time on an otherwise idle B200. Batch invariance is bitwise and covers +three checks: every row computed alone vs inside full batches of 64, 257 and 2048 rows; the full +8192-timestep batch vs sub-batches that together cover every row; and a dense batch-size sweep. A +"no" means that at least one row, sub-batch or gradient differed. [`timestep_sinusoid.json`](../../reports/experiments/h3-prior-art-b200/timestep_sinusoid.json) +was written from a clean tree at `7a82917` by + +```bash +python tools/validation/models/h3_prior_art.py --op timestep_sinusoid --out reports/experiments/h3-prior-art-b200/timestep_sinusoid.json +python tools/validation/models/plot_h3_prior_art.py reports/experiments/h3-prior-art-b200/timestep_sinusoid.json +``` + +Libraries that do not import are skipped and recorded as unavailable in the report. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_timestep_sinusoid.py -v # operator +python -m pytest tests/models/minimax_h3/test_h3_conditioning_e2e.py -v # end to end, pinned weights +python tools/validation/operators/check_operator.py --op timestep_sinusoid_h3 --candidate cuda --device cuda \ + --dtype fp32 --batch 7 --check-grad +python tools/validation/models/h3_evidence.py --op timestep_sinusoid_h3 \ + --out reports/experiments/h3-timestep-sinusoid-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-timestep-sinusoid-b200/report.json +``` + +## Known Limitations + +- There is no ROCm kernel. ROCm dispatches the PyTorch reference. +- `max_period` is fixed at 10000, the only value H3 uses. diff --git a/docs/operators/h3-tp-adaln-3mod.md b/docs/operators/h3-tp-adaln-3mod.md new file mode 100644 index 000000000..c348bce35 --- /dev/null +++ b/docs/operators/h3-tp-adaln-3mod.md @@ -0,0 +1,113 @@ +# MiniMax-H3 Tensor-Parallel AdaLN Projection + +## Summary + +`tp_adaln_3mod` shards the [AdaLN projection](h3-adaln-projection.md) (`2688 -> 96768`, 520 MB of +BF16 weight per block) across tensor-parallel ranks (RFC #420, WS2). Every rank ends with the +**same bytes as the single-GPU (WS1) op**: the six modulation tensors, the `3T` modality rows, +`d_temb`, and its shard of `dW`/`db`. + +## Entry Point + +```python +from rl_engine.distributed.algorithms.collectives import DeterministicCollective +from rl_engine.backends.cuda.model_specific.minimax_h3.tp_adaln_projection import ( + H3TPAdaLNProjectionCudaOp, + shard_adaln_projection, +) + +collective = DeterministicCollective(group=tp_group, device=local_rank) +op = H3TPAdaLNProjectionCudaOp(collective, n_total=96768) +w_shard, b_shard = shard_adaln_projection(weight, bias, collective.world_size, collective.rank) +shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = op(temb, w_shard, b_shard) +op.readback() # tp, rank, owned columns/slots, collective backend, kernel ids, fallback +``` + +## Ownership + +Rank `r` of `tp` owns the contiguous output columns `[r * N / tp, (r + 1) * N / tp)` of the +`N = 3 x 6 x 5376` table, which are the same rows of `adaln_proj.linear.weight` and `bias`. +`N / tp` must be a multiple of 64, the `d_input` chunk of contract `h3-det-linear-v1`, so a shard +never splits a WS1 partial sum. The shard constructor accepts positive TP sizes satisfying +this alignment, including 1, 2, 3, 4, 6, 7, 8 and 9 (TP=9 gives 10752 columns per rank). +The `DeterministicCollective` backend supports only 1, 2, 4 and 8 ranks; shard alignment +alone does not imply backend support. + +| TP | Columns per rank | Rank 0 owns | +| --- | --- | --- | +| 2 | 48384 | video: all six chunks; text: shift/scale/gate_msa | +| 4 | 24192 | video: shift_msa .. shift_mlp, half of scale_mlp | +| 8 | 12096 | video: shift_msa, scale_msa, 1344 channels of gate_msa | + +`AdaLNColumnShard.slots()` lists each rank's `(modality, chunk, h_begin, h_end)` pieces. The +all-gathered table is the WS1 table, so the six tensors and the `t * 3 + modality` rows are views +of it exactly as in WS1. Nothing downstream needs to know the TP size. + +## Numerics + +Every floating-point operation is the WS1 kernel itself. The only collective is a rank-ordered +all-gather, which is a copy. No collective reduction arithmetic is used. + +- **Forward.** An output column depends only on `silu(temb)` and its own weight row, so the local + tensor-core GEMV yields exactly the WS1 columns. The all-gather rebuilds the `(T, N)` table. +- **`dW`, `db`.** These are rows of the shard, computed locally with the WS1 ascending-`t` fold. +- **`d_temb`.** WS1 sums `N` in 64-row chunks and then left-folds the chunk partials in ascending + order. Each rank computes the partials of its own chunks + (`h3_det_linear_backward_input_partials`), the all-gather puts them in global chunk order, and + every rank runs the WS1 fold (`h3_det_linear_fold_chunks`) and the FP32 SiLU VJP. + +The backward reads its own columns of the table gradient. It therefore requires that gradient to +be identical on every TP rank, as it is when every rank applies the modulation to the same rows. + +## Communication + +| Direction | Payload per rank | At T = 3, TP = 8 | +| --- | --- | --- | +| forward | `T x N / tp` BF16 | 72.6 KB | +| backward | `(N / 64 / tp) x T x 2688` FP32 | 6.1 MB | + +`DeterministicCollective` copies all-gathers whose output is at most 256 KiB with a single-block +byte loop, at about 4 µs per KiB on B200. `ws2_comm.gather_rows` pads such shards just past that +size, where the multi-block path takes about 45 µs. This matters for the forward at small `T`. + +## Evidence + +![tp_adaln_3mod on 8 x B200: byte equality and time](../../reports/experiments/h3-tp-adaln-3mod-b200/figure.png) + +The report is written by `tools/validation/models/h3_ws2_evidence.py` from a clean tree at commit `4947996`, with +real NCCL processes on one 8 x B200 node, one GPU per rank, on the pinned block-0 weights. For +TP 1, 2, 4 and 8 and T = 1..4, every rank's table and `d_temb`, and its `dW`/`db` shard, are +byte-equal to WS1 computed on the same GPU. + +| T | | WS1 (1 GPU) | TP2 | TP4 | TP8 | +| --- | --- | --- | --- | --- | --- | +| 1 | forward | 0.10 ms | 0.12 ms | 0.11 ms | 0.12 ms | +| 1 | forward + backward | 1.90 ms | 1.49 ms | 1.19 ms | 1.04 ms | +| 3 | forward | 0.10 ms | 0.11 ms | 0.11 ms | 0.12 ms | +| 3 | forward + backward | 2.25 ms | 1.77 ms | 1.48 ms | 1.32 ms | +| 4 | forward + backward | 2.09 ms | 1.83 ms | 1.57 ms | 1.45 ms | + +Times are for the slowest rank. The WS1 forward already streams the 520 MB weight at about +5 TB/s, so the TP forward is bounded by the all-gather. `DeterministicCollective` synchronises the +host on every call, which costs about 0.1 ms. The backward is dominated by the shard-local `dW` +and gains 1.7-1.8x at TP8. Each rank also holds only `1 / tp` of the weight and its gradient. + +## Tests + +```bash +export RL_KERNEL_H3_WEIGHTS= +python -m pytest tests/models/minimax_h3/test_h3_tp_adaln.py -v # ownership, shard arithmetic, NCCL TP2/4/8 +python tools/validation/models/h3_ws2_evidence.py --op tp_adaln_3mod --worlds 1,2,4,8 \ + --out reports/experiments/h3-tp-adaln-3mod-b200/report.json +python tools/validation/models/plot_h3_evidence.py reports/experiments/h3-tp-adaln-3mod-b200/report.json +``` + +The NCCL tests need as many GPUs as ranks and skip otherwise. The shard-arithmetic tests check +the same property on one GPU: each rank's columns and chunk partials equal the matching slice of +the WS1 call, and the fold of the gathered partials equals WS1's `d_input`. + +## Known Limitations + +- CUDA only. `DeterministicCollective` supports world sizes 1, 2, 4 and 8. +- `norm_out.linear` (`2688 -> 10752`) is not sharded here. It is small and stays replicated. +- There is no ROCm path. diff --git a/docs/usage/evidence/h3-one-block-b200/figure.png b/docs/usage/evidence/h3-one-block-b200/figure.png new file mode 100644 index 000000000..86cb58af1 Binary files /dev/null and b/docs/usage/evidence/h3-one-block-b200/figure.png differ diff --git a/docs/usage/evidence/h3-one-block-b200/report.json b/docs/usage/evidence/h3-one-block-b200/report.json new file mode 100644 index 000000000..20a7f4bcd --- /dev/null +++ b/docs/usage/evidence/h3-one-block-b200/report.json @@ -0,0 +1,1875 @@ +{ + "kind": "h3_one_block_replay", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "weights_sha256": { + "diffusion_pytorch_model-00001-of-00014.safetensors": { + "sha256": "2d847200c45c09dd7f973c1b096663068408ef851ee0b3711d059b6dc5dcd028", + "size_bytes": 4825958704 + } + }, + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "41e429327996d5f2fa7485561d158d66f9308e70", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "nodes": [ + { + "node": "adaln_projection", + "inputs": [ + "temb" + ], + "rfc_row": "adaln_projection_3mod", + "status": "row", + "backend": "H3AdaLNProjectionCudaOp" + }, + { + "node": "norm1", + "inputs": [ + "hidden", + "adaln_projection" + ], + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated" + }, + { + "node": "q_proj", + "inputs": [ + "norm1" + ], + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "k_proj", + "inputs": [ + "norm1" + ], + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "v_proj", + "inputs": [ + "norm1" + ], + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "q_norm", + "inputs": [ + "q_proj" + ], + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward" + }, + { + "node": "k_norm", + "inputs": [ + "k_proj" + ], + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward" + }, + { + "node": "rope_q", + "inputs": [ + "q_norm" + ], + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary" + }, + { + "node": "rope_k", + "inputs": [ + "k_norm" + ], + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary" + }, + { + "node": "attention", + "inputs": [ + "rope_q", + "rope_k", + "v_proj" + ], + "rfc_row": "h3_full_attention", + "status": "interim", + "backend": "DeterministicAttentionOp(causal=False)" + }, + { + "node": "o_proj", + "inputs": [ + "attention" + ], + "rfc_row": "h3_attention_o_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "residual_attn", + "inputs": [ + "hidden", + "o_proj", + "adaln_projection" + ], + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp" + }, + { + "node": "norm2", + "inputs": [ + "residual_attn", + "adaln_projection" + ], + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated" + }, + { + "node": "ffn_gate_up", + "inputs": [ + "norm2" + ], + "rfc_row": "h3_ffn_gate_up_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "swiglu", + "inputs": [ + "ffn_gate_up" + ], + "rfc_row": "h3_swiglu", + "status": "interim", + "backend": "SwiGLUCudaOp" + }, + { + "node": "ffn_down", + "inputs": [ + "swiglu" + ], + "rfc_row": "h3_ffn_down_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + }, + { + "node": "residual_mlp", + "inputs": [ + "residual_attn", + "ffn_down", + "adaln_projection" + ], + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp" + } + ], + "forward_cases": [ + { + "layout": "tiny", + "seq_len": 264, + "seed": 0, + "backends": { + "projection": "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection.H3AdaLNProjectionCudaOp", + "norm": "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp", + "gate": "rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual.H3GateResidualCudaOp", + "gemm": "rl_engine.backends.shared.triton.gemm.det_gemm.TritonDetGemmOp", + "attention": "rl_engine.backends.cuda.attention.deterministic_attn.DeterministicAttentionOp", + "swiglu": "rl_engine.backends.cuda.activation.swiglu.SwiGLUCudaOp" + }, + "nodes": [ + { + "node": "adaln_projection", + "rfc_row": "adaln_projection_3mod", + "status": "row", + "backend": "H3AdaLNProjectionCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "first_differing_row": -1 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625 + }, + "candidate_rel_err_vs_golden": 0.00246217805536275, + "provider_rel_err_vs_golden": 0.00246217805536275, + "candidate_sha256": "68f9c15a17db165510db500aedb040d0faa8f8b67c529b73eb9f3dcd8c3e525b" + }, + { + "node": "norm1", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0087890625, + "first_differing_row": 32 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005880745783000111, + "provider_rel_err_vs_golden": 0.005880745783000111, + "candidate_sha256": "ea1c6f4c18d82295c350c586fa038e9e1c867f7908dd222f3818e4c991752ff8" + }, + { + "node": "q_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25 + }, + "candidate_rel_err_vs_golden": 0.0026542556019516343, + "provider_rel_err_vs_golden": 0.0026542556019516343, + "candidate_sha256": "caa29aa7557916f97ffcca23acee8d2d00be9350e5a1c7e562ab199b2efc1e84" + }, + { + "node": "k_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.0027707960565066404, + "provider_rel_err_vs_golden": 0.0027707960565066404, + "candidate_sha256": "3bbed972cc826de83b70684a9d4161e1cd4f6a9c2b9e4d9814c5cd1a64b76600" + }, + { + "node": "v_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625 + }, + "candidate_rel_err_vs_golden": 0.003925822833984887, + "provider_rel_err_vs_golden": 0.003925822833984887, + "candidate_sha256": "0cc7c296450b743d1958eed5e0a17275b3f2c143b3ab5166400edad7bc6d79bb" + }, + { + "node": "q_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005536327574840286, + "provider_rel_err_vs_golden": 0.005536327574840286, + "candidate_sha256": "8e2244615dfbc87506392629ca682fdc6eb6f9280f0323bede3bc41484b539f6" + }, + { + "node": "k_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.004772805574992097, + "provider_rel_err_vs_golden": 0.004772805574992097, + "candidate_sha256": "5db66dbaf307aa503ed36484d4846dba4f33c5042a8f64e94144ce58b65b12e9" + }, + { + "node": "rope_q", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.0068777999680434765, + "provider_rel_err_vs_golden": 0.0068777999680434765, + "candidate_sha256": "6d7d9f7b0d3ae0b8a504249fa855320d6744742c65feea95b58b1ef4fc593eb2" + }, + { + "node": "rope_k", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.007537748803240922, + "provider_rel_err_vs_golden": 0.007537748803240922, + "candidate_sha256": "2fffe79bf75b348798bb739fa268fca67091afafb1e1f3cea548288769101811" + }, + { + "node": "attention", + "rfc_row": "h3_full_attention", + "status": "interim", + "backend": "DeterministicAttentionOp(causal=False)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625 + }, + "candidate_rel_err_vs_golden": 0.0083336167695551, + "provider_rel_err_vs_golden": 0.0083336167695551, + "candidate_sha256": "32b27348e7fe94d770fa79fc8e59c485be7cefeed05794bf29635f4e145e652a" + }, + { + "node": "o_proj", + "rfc_row": "h3_attention_o_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.5 + }, + "candidate_rel_err_vs_golden": 0.003780642806114357, + "provider_rel_err_vs_golden": 0.003780642806114357, + "candidate_sha256": "1e20c6f980d2b66055fe7f4af469b10f05949a29d08d5f0f55edb9fd03b5a016" + }, + { + "node": "residual_attn", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 4.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005879187649263115, + "provider_rel_err_vs_golden": 0.005879187649263115, + "candidate_sha256": "0c4eb6f34c400e6f6d0b5b5e75ec87456902a5c914815ef4c63979b54344d93d" + }, + { + "node": "norm2", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.012571502416896373, + "provider_rel_err_vs_golden": 0.012571502416896373, + "candidate_sha256": "a422291adc6ba08185e3a0a64989ce061b7d91f261e84d7906124a389fa44a01" + }, + { + "node": "ffn_gate_up", + "rfc_row": "h3_ffn_gate_up_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625 + }, + "candidate_rel_err_vs_golden": 0.009016771723469647, + "provider_rel_err_vs_golden": 0.009738541525197427, + "candidate_sha256": "eb4f1cab146f89a5c6dbfee54c8d9e14f4a78b7366b8f25829799e952e367221" + }, + { + "node": "swiglu", + "rfc_row": "h3_swiglu", + "status": "interim", + "backend": "SwiGLUCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 1.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.5 + }, + "candidate_rel_err_vs_golden": 0.013677000528527118, + "provider_rel_err_vs_golden": 0.01192762915547773, + "candidate_sha256": "04b5c030f9c8ef663a9f51004627b0f6a540d7186b3008b32714bff5877d4f50" + }, + { + "node": "ffn_down", + "rfc_row": "h3_ffn_down_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 16.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0 + }, + "candidate_rel_err_vs_golden": 0.004798385019493273, + "provider_rel_err_vs_golden": 0.004798385019493273, + "candidate_sha256": "f4c9bb9f93792b7ca744915621e097325783910d9466a43426ddc7ac1f3420b9" + }, + { + "node": "residual_mlp", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 64.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005205170904616623, + "provider_rel_err_vs_golden": 0.005205170904616623, + "candidate_sha256": "f94be57b6a6e580cf709fce6ad40bde5c4386aa7f0ed3b83707159b8c42ec9bc" + } + ], + "first_drift": "adaln_projection", + "first_isolated_drift": "adaln_projection", + "first_repeat_drift": null, + "first_batch_drift": null, + "output": { + "candidate_rel_err_vs_golden": 0.005205170904616623, + "provider_rel_err_vs_golden": 0.005205170904616623, + "candidate_vs_provider_max_abs": 64.0 + }, + "permutation_probe": { + "bitwise_equal_after_permutation": false, + "max_abs": 2.0, + "differing_rows": 79 + }, + "timing_ms": { + "candidate_forward": 7.439708977472037, + "provider_forward": 0.6147729873191565 + } + }, + { + "layout": "small", + "seq_len": 925, + "seed": 0, + "backends": { + "projection": "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection.H3AdaLNProjectionCudaOp", + "norm": "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp", + "gate": "rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual.H3GateResidualCudaOp", + "gemm": "rl_engine.backends.shared.triton.gemm.det_gemm.TritonDetGemmOp", + "attention": "rl_engine.backends.cuda.attention.deterministic_attn.DeterministicAttentionOp", + "swiglu": "rl_engine.backends.cuda.activation.swiglu.SwiGLUCudaOp" + }, + "nodes": [ + { + "node": "adaln_projection", + "rfc_row": "adaln_projection_3mod", + "status": "row", + "backend": "H3AdaLNProjectionCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "first_differing_row": -1 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625 + }, + "candidate_rel_err_vs_golden": 0.00246217805536275, + "provider_rel_err_vs_golden": 0.00246217805536275, + "candidate_sha256": "68f9c15a17db165510db500aedb040d0faa8f8b67c529b73eb9f3dcd8c3e525b" + }, + { + "node": "norm1", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "first_differing_row": 77 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.006146196627110127, + "provider_rel_err_vs_golden": 0.006146196627110127, + "candidate_sha256": "6757e07bd6c2c8a5ce53cf308b444fcbdb8314b29fc38f3bf9a0fec07c887e3a" + }, + { + "node": "q_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25 + }, + "candidate_rel_err_vs_golden": 0.0029038112813761937, + "provider_rel_err_vs_golden": 0.0029038112813761937, + "candidate_sha256": "d958bf5d52ca54544694b8b85bb83105810a388c5aff8376929fbcd03a04b840" + }, + { + "node": "k_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.003283722873092741, + "provider_rel_err_vs_golden": 0.003283722873092741, + "candidate_sha256": "57140b54d246601eb859ae46ca9a935ac20447cb07ecda1f63bb4760664e630a" + }, + { + "node": "v_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625 + }, + "candidate_rel_err_vs_golden": 0.00430988563576781, + "provider_rel_err_vs_golden": 0.00430988563576781, + "candidate_sha256": "a0c247fbd0bdea19f440ccd0bad518297c03aeeda32aa9d36055cc4d3dd7a1f6" + }, + { + "node": "q_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.004885448038371641, + "provider_rel_err_vs_golden": 0.004885448038371641, + "candidate_sha256": "7798faea2b91461666a1bf951d7f035f6f237482313f83ca72a84a4ff46bb33f" + }, + { + "node": "k_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005799367284411749, + "provider_rel_err_vs_golden": 0.005799367284411749, + "candidate_sha256": "8847483fb44d687353d588a8f34e79e3681f7e5eaa7b1d61cc1a1a90d2eca8a0" + }, + { + "node": "rope_q", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008034768012785918, + "provider_rel_err_vs_golden": 0.008034768012785918, + "candidate_sha256": "1e6f9c0f0e1f1a93736fed847964760bceddc7c5f8ffee316586ea7edf080ae1" + }, + { + "node": "rope_k", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008003912133033995, + "provider_rel_err_vs_golden": 0.008003912133033995, + "candidate_sha256": "08f1bae87e8e03f52e9baff203daceee7c6689b39d98d7e71d5c2386125f15c6" + }, + { + "node": "attention", + "rfc_row": "h3_full_attention", + "status": "interim", + "backend": "DeterministicAttentionOp(causal=False)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.01508755753795434, + "provider_rel_err_vs_golden": 0.01508755753795434, + "candidate_sha256": "ce752e0660f62458877036eb57ee0962fadf2d461a66335d4cf2cbd5fa535b1e" + }, + { + "node": "o_proj", + "rfc_row": "h3_attention_o_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0 + }, + "candidate_rel_err_vs_golden": 0.0033698312381439145, + "provider_rel_err_vs_golden": 0.0033698312381439145, + "candidate_sha256": "fb87b744c614e9031b3d0bedaa17a047aecdfa73f6879bc1ad226027c5fbc830" + }, + { + "node": "residual_attn", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008111918367374191, + "provider_rel_err_vs_golden": 0.008111918367374191, + "candidate_sha256": "5ded0e906f14fe1b8c07c8ab1c0a64c4e54b55016280b853ce66e847bbb30496" + }, + { + "node": "norm2", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.01274206642706212, + "provider_rel_err_vs_golden": 0.01274206642706212, + "candidate_sha256": "ef9c8d46b1ab82858d0e52b9738d15bb5804583a5c2521881faaf8a60dd0d900" + }, + { + "node": "ffn_gate_up", + "rfc_row": "h3_ffn_gate_up_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.140625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.008656912666959045, + "provider_rel_err_vs_golden": 0.008857528206022237, + "candidate_sha256": "4a703fccc5383ea70c3d0ef7ceb5140302994c20fd8d629e1c7f5f1267ffd0bb" + }, + { + "node": "swiglu", + "rfc_row": "h3_swiglu", + "status": "interim", + "backend": "SwiGLUCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 1.5, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.5 + }, + "candidate_rel_err_vs_golden": 0.014360147174659211, + "provider_rel_err_vs_golden": 0.014360147174659211, + "candidate_sha256": "ca07f060d2a24bee331f60fae51cfcf52a69ecb1aabd6b5e696d732713ec13ca" + }, + { + "node": "ffn_down", + "rfc_row": "h3_ffn_down_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 16.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 4.0 + }, + "candidate_rel_err_vs_golden": 0.007665430184168193, + "provider_rel_err_vs_golden": 0.006803410809063402, + "candidate_sha256": "341dad90d42a52a36a0afbe490699d0a50020f1a0ec1c7c77d8938b44a5731fb" + }, + { + "node": "residual_mlp", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 64.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.007871560571379407, + "provider_rel_err_vs_golden": 0.007871560571379407, + "candidate_sha256": "7f6e08e2998a410c9a115092f16ae4d95c330e7b88afba9f80c4bd9367b7417a" + } + ], + "first_drift": "adaln_projection", + "first_isolated_drift": "adaln_projection", + "first_repeat_drift": null, + "first_batch_drift": null, + "output": { + "candidate_rel_err_vs_golden": 0.007871560571379407, + "provider_rel_err_vs_golden": 0.007871560571379407, + "candidate_vs_provider_max_abs": 64.0 + }, + "permutation_probe": { + "bitwise_equal_after_permutation": false, + "max_abs": 8.0, + "differing_rows": 450 + }, + "timing_ms": { + "candidate_forward": 30.122189986286685, + "provider_forward": 1.124817004892975 + } + }, + { + "layout": "medium", + "seq_len": 3160, + "seed": 0, + "backends": { + "projection": "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection.H3AdaLNProjectionCudaOp", + "norm": "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp", + "gate": "rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual.H3GateResidualCudaOp", + "gemm": "rl_engine.backends.shared.triton.gemm.det_gemm.TritonDetGemmOp", + "attention": "rl_engine.backends.cuda.attention.deterministic_attn.DeterministicAttentionOp", + "swiglu": "rl_engine.backends.cuda.activation.swiglu.SwiGLUCudaOp" + }, + "nodes": [ + { + "node": "adaln_projection", + "rfc_row": "adaln_projection_3mod", + "status": "row", + "backend": "H3AdaLNProjectionCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "first_differing_row": -1 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625 + }, + "candidate_rel_err_vs_golden": 0.00246217805536275, + "provider_rel_err_vs_golden": 0.00246217805536275, + "candidate_sha256": "68f9c15a17db165510db500aedb040d0faa8f8b67c529b73eb9f3dcd8c3e525b" + }, + { + "node": "norm1", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01171875, + "first_differing_row": 120 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.006792479568272936, + "provider_rel_err_vs_golden": 0.006792479568272936, + "candidate_sha256": "6989c76ea21e7f0a0d5c659b3f3fd0b84a0f990c0ac9c7ae3ebf61af0643c65a" + }, + { + "node": "q_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25 + }, + "candidate_rel_err_vs_golden": 0.0028271579867662055, + "provider_rel_err_vs_golden": 0.0028271579867662055, + "candidate_sha256": "d954f84762d25e78aecdcb79ccd99c994296437fbfbc37d2ec578f55ccc0f280" + }, + { + "node": "k_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.25 + }, + "candidate_rel_err_vs_golden": 0.003283722873092741, + "provider_rel_err_vs_golden": 0.003283722873092741, + "candidate_sha256": "eba776772bf3c5e7de31cf613be33ffa3ce309513e4ac276fb19df73dc5db006" + }, + { + "node": "v_proj", + "rfc_row": "h3_qkv_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.00430988563576781, + "provider_rel_err_vs_golden": 0.00430988563576781, + "candidate_sha256": "ed928de4d7846e2364135e6c7d4db80fc8a2f6f1eeb8cbcbab1c3759a5858f7b" + }, + { + "node": "q_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.00552281985076463, + "provider_rel_err_vs_golden": 0.00552281985076463, + "candidate_sha256": "0a46cc880e887a2048ebc03a3b24590b1ebd1fca7d34f5ebc5134c1f7435bcf3" + }, + { + "node": "k_norm", + "rfc_row": "h3_qk_rmsnorm_d128", + "status": "interim", + "backend": "H3RMSNormCudaOp.forward", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.005531504705471249, + "provider_rel_err_vs_golden": 0.005531504705471249, + "candidate_sha256": "7500467f678d2fd82333a3e641c1608b5dc7cbedcaa6b380b361efee63aec875" + }, + { + "node": "rope_q", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008244304699552833, + "provider_rel_err_vs_golden": 0.008244304699552833, + "candidate_sha256": "c1638aca15cb9d18cd51ffb2ea1959b946984135babd90f2398ec804321d2015" + }, + { + "node": "rope_k", + "rfc_row": "h3_mm_rope_3axis_partial", + "status": "reference", + "backend": "provider_apply_rotary", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.09375, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.009406142205096895, + "provider_rel_err_vs_golden": 0.009406142205096895, + "candidate_sha256": "29ac16d0d67970fdd0c8d99db62767acfe199d891d56d82dc82d011aaf7bd9d0" + }, + { + "node": "attention", + "rfc_row": "h3_full_attention", + "status": "interim", + "backend": "DeterministicAttentionOp(causal=False)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625 + }, + "candidate_rel_err_vs_golden": 0.01760013142702786, + "provider_rel_err_vs_golden": 0.014630552458297459, + "candidate_sha256": "cf086b941114d523587b8a4243ee34085464e69c93901e272c76945bce1472b6" + }, + { + "node": "o_proj", + "rfc_row": "h3_attention_o_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0 + }, + "candidate_rel_err_vs_golden": 0.003938049145471708, + "provider_rel_err_vs_golden": 0.0033432092275601103, + "candidate_sha256": "08f5333667b8d6cf0f93096a52b77dd3662b032f89b701c63fe5bc860f9aaaf0" + }, + { + "node": "residual_attn", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 4.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008632717009347793, + "provider_rel_err_vs_golden": 0.008632717009347793, + "candidate_sha256": "97d5218779c288296d89826e22e4cf1263dd4aa2ca11882de0c5fbc930266b70" + }, + { + "node": "norm2", + "rfc_row": "h3_rmsnorm", + "status": "row", + "backend": "H3RMSNormCudaOp.forward_modulated", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.046875, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.014276226801909717, + "provider_rel_err_vs_golden": 0.014276226801909717, + "candidate_sha256": "79869eb1924b37222bf608af880f848e3092462415df4414377ed11302b255ba" + }, + { + "node": "ffn_gate_up", + "rfc_row": "h3_ffn_gate_up_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.1875, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.125 + }, + "candidate_rel_err_vs_golden": 0.009758666244112313, + "provider_rel_err_vs_golden": 0.009758666244112313, + "candidate_sha256": "2cc1191d175f5e594b8a6ac7ab3f15b0232d3b862a8fe2e5dd094c38edb0c3ef" + }, + { + "node": "swiglu", + "rfc_row": "h3_swiglu", + "status": "interim", + "backend": "SwiGLUCudaOp", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 2.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.5 + }, + "candidate_rel_err_vs_golden": 0.013815450163525256, + "provider_rel_err_vs_golden": 0.014338350515051697, + "candidate_sha256": "b2893d7362bb9c9ffb3ad9b8db439a3f72161f2bbcdba1465fa28112060d9c23" + }, + { + "node": "ffn_down", + "rfc_row": "h3_ffn_down_gemm", + "status": "interim", + "backend": "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)", + "provider_bitwise_isolated_promised": false, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 16.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 16.0 + }, + "candidate_rel_err_vs_golden": 0.007855981935113145, + "provider_rel_err_vs_golden": 0.008322311373457961, + "candidate_sha256": "0be4ddb65e1fa0aa070ca6cd7ce3045b2c75b84a2565d0607e1c092a942313cf" + }, + { + "node": "residual_mlp", + "rfc_row": "adaln_gate_residual", + "status": "row", + "backend": "H3GateResidualCudaOp", + "provider_bitwise_isolated_promised": true, + "repeat_bitwise_equal": true, + "batch_row_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 64.0, + "first_differing_row": 0 + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0 + }, + "candidate_rel_err_vs_golden": 0.008972776288585085, + "provider_rel_err_vs_golden": 0.008972776288585085, + "candidate_sha256": "6ad8b3fb011f0371de6cb5e63e133d2bdfa8e2644eb6fed3c7ce4b0a60a59f16" + } + ], + "first_drift": "adaln_projection", + "first_isolated_drift": "adaln_projection", + "first_repeat_drift": null, + "first_batch_drift": null, + "output": { + "candidate_rel_err_vs_golden": 0.008972776288585085, + "provider_rel_err_vs_golden": 0.008972776288585085, + "candidate_vs_provider_max_abs": 64.0 + }, + "permutation_probe": { + "bitwise_equal_after_permutation": false, + "max_abs": 8.0, + "differing_rows": 2199 + }, + "timing_ms": { + "candidate_forward": 225.18538997974247, + "provider_forward": 3.804815001785755 + } + } + ], + "backward_cases": [ + { + "layout": "tiny", + "seq_len": 264, + "seed": 0, + "leaves": { + "adaln_proj.linear.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.01180618547406291, + "provider_rel_err_vs_golden": 0.02146984829363087 + }, + "adaln_proj.linear.bias": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.00719525178157443, + "provider_rel_err_vs_golden": 0.01182990495831936 + }, + "norm1.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.009806818937881757, + "provider_rel_err_vs_golden": 0.011246836628335638 + }, + "attn.to_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.021136961430992755, + "provider_rel_err_vs_golden": 0.01848097081774083 + }, + "attn.to_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.012966444280077611, + "provider_rel_err_vs_golden": 0.013865612206340512 + }, + "attn.to_v.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.009565037054795875, + "provider_rel_err_vs_golden": 0.008223319017080894 + }, + "attn.norm_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.03398653906442352, + "provider_rel_err_vs_golden": 0.02988808485881452 + }, + "attn.norm_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.024343253203742525, + "provider_rel_err_vs_golden": 0.021079043911823313 + }, + "attn.to_out.0.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.00622458042676816, + "provider_rel_err_vs_golden": 0.006710293163390872 + }, + "norm2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.015949579365854474, + "provider_rel_err_vs_golden": 0.015949579365854474 + }, + "ff.net.0.proj.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.016500518451181855, + "provider_rel_err_vs_golden": 0.016500518451181855 + }, + "ff.net.2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.026654513525735446, + "provider_rel_err_vs_golden": 0.026654513525735446 + }, + "hidden": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.012426632450983136, + "provider_rel_err_vs_golden": 0.012460023186919867, + "batch_row_bitwise_equal": true + }, + "temb": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.0099260042974754, + "provider_rel_err_vs_golden": 0.009738359099414324 + } + }, + "nodes": [ + { + "node": "residual_mlp", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_down", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "swiglu", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_gate_up", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm2", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "residual_attn", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "o_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "attention", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_k", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_q", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "v_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm1", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "adaln_projection", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": null + } + ], + "first_grad_repeat_drift": null, + "first_grad_batch_drift": null, + "timing_ms": { + "candidate_forward_backward": 22.65317298588343, + "provider_forward_backward": 2.773943997453898 + } + }, + { + "layout": "small", + "seq_len": 925, + "seed": 0, + "leaves": { + "adaln_proj.linear.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.021554155200271704, + "provider_rel_err_vs_golden": 0.034999791088826775 + }, + "adaln_proj.linear.bias": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.01963672978398978, + "provider_rel_err_vs_golden": 0.03383649401393024 + }, + "norm1.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.018085877508142802, + "provider_rel_err_vs_golden": 0.018085877508142802 + }, + "attn.to_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.013599306658661192, + "provider_rel_err_vs_golden": 0.01387979041919243 + }, + "attn.to_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.019453065667697233, + "provider_rel_err_vs_golden": 0.020302731870153884 + }, + "attn.to_v.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.012148514788997115, + "provider_rel_err_vs_golden": 0.013811853109547911 + }, + "attn.norm_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.013203050956559442, + "provider_rel_err_vs_golden": 0.0244166526502349 + }, + "attn.norm_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.016342136638754914, + "provider_rel_err_vs_golden": 0.021135696931241458 + }, + "attn.to_out.0.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.008114640734095811, + "provider_rel_err_vs_golden": 0.009028884858896587 + }, + "norm2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.014249048918878491, + "provider_rel_err_vs_golden": 0.02784561376137672 + }, + "ff.net.0.proj.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.014849208055237118, + "provider_rel_err_vs_golden": 0.01293128762801526 + }, + "ff.net.2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.007383423169385204, + "provider_rel_err_vs_golden": 0.00660931488762284 + }, + "hidden": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.017912377413917616, + "provider_rel_err_vs_golden": 0.021024368694072988, + "batch_row_bitwise_equal": true + }, + "temb": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.011320712818158819, + "provider_rel_err_vs_golden": 0.015373420760022231 + } + }, + "nodes": [ + { + "node": "residual_mlp", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_down", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "swiglu", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_gate_up", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm2", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "residual_attn", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "o_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "attention", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_k", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_q", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "v_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm1", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "adaln_projection", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": null + } + ], + "first_grad_repeat_drift": null, + "first_grad_batch_drift": null, + "timing_ms": { + "candidate_forward_backward": 81.66394199361093, + "provider_forward_backward": 4.987127991626039 + } + }, + { + "layout": "medium", + "seq_len": 3160, + "seed": 0, + "leaves": { + "adaln_proj.linear.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.0067601458928464005, + "provider_rel_err_vs_golden": 0.045201668088719865 + }, + "adaln_proj.linear.bias": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.0071676229870052495, + "provider_rel_err_vs_golden": 0.04860427948029884 + }, + "norm1.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.00807488848365698, + "provider_rel_err_vs_golden": 0.00807488848365698 + }, + "attn.to_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.013988988631677083, + "provider_rel_err_vs_golden": 0.013988988631677083 + }, + "attn.to_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.020358272407898427, + "provider_rel_err_vs_golden": 0.01588018752234606 + }, + "attn.to_v.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.013133613529357416, + "provider_rel_err_vs_golden": 0.013503072539846854 + }, + "attn.norm_q.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.00947842982880085, + "provider_rel_err_vs_golden": 0.005505274396843249 + }, + "attn.norm_k.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.009501336513049831, + "provider_rel_err_vs_golden": 0.00797009449782898 + }, + "attn.to_out.0.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.00966903564860318, + "provider_rel_err_vs_golden": 0.010693027840662716 + }, + "norm2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.012098893170937968, + "provider_rel_err_vs_golden": 0.012098893170937968 + }, + "ff.net.0.proj.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.008304478641992876, + "provider_rel_err_vs_golden": 0.009736491756659858 + }, + "ff.net.2.weight": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.015244679882575392, + "provider_rel_err_vs_golden": 0.015244679882575392 + }, + "hidden": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.01881164387407051, + "provider_rel_err_vs_golden": 0.020717783234680102, + "batch_row_bitwise_equal": true + }, + "temb": { + "repeat_bitwise_equal": true, + "candidate_rel_err_vs_golden": 0.008158318242825988, + "provider_rel_err_vs_golden": 0.0566394029843732 + } + }, + "nodes": [ + { + "node": "residual_mlp", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_down", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "swiglu", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "ffn_gate_up", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm2", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "residual_attn", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "o_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "attention", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_k", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "rope_q", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_norm", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "v_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "k_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "q_proj", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "norm1", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": true + }, + { + "node": "adaln_projection", + "grad_repeat_bitwise_equal": true, + "grad_batch_row_bitwise_equal": null + } + ], + "first_grad_repeat_drift": null, + "first_grad_batch_drift": null, + "timing_ms": { + "candidate_forward_backward": 577.440707013011, + "provider_forward_backward": 11.53312501264736 + } + } + ] +} diff --git a/reports/experiments/h3-adaln-gate-residual-b200/figure.png b/reports/experiments/h3-adaln-gate-residual-b200/figure.png new file mode 100644 index 000000000..a2d4e2a59 Binary files /dev/null and b/reports/experiments/h3-adaln-gate-residual-b200/figure.png differ diff --git a/reports/experiments/h3-adaln-gate-residual-b200/report.json b/reports/experiments/h3-adaln-gate-residual-b200/report.json new file mode 100644 index 000000000..750d13ecf --- /dev/null +++ b/reports/experiments/h3-adaln-gate-residual-b200/report.json @@ -0,0 +1,6187 @@ +{ + "kind": "h3_operator_report", + "op": "adaln_gate_residual", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "weight_source": "pinned_checkpoint", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "bde1d8d97975a9b5ad32b7f0f8486e7f6815638b", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "forward_bitwise_vs_diffusers": { + "bfloat16": true, + "float16": true, + "float32": true + }, + "rows_batch_invariant": true, + "backward": { + "cuda": { + "repeat_bitwise_equal": true, + "rel_error": { + "d_residual": 0.0, + "d_sublayer": 0.002054442732408834, + "d_gate": 0.003085325221208727 + } + }, + "provider": { + "repeat_bitwise_equal": false, + "rel_error": { + "d_residual": 0.0, + "d_sublayer": 0.002054442732408834, + "d_gate": 0.0407378032311299 + } + } + } + }, + "perf": [ + { + "op": "adaln_gate_residual", + "case": "S=4097", + "backend": "H3GateResidualCudaOp", + "bytes": 132152832, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 57.37600103020668, + 76.22399926185608, + 49.47200044989586, + 76.60800218582153, + 49.75999891757965, + 75.83999633789062, + 48.31999912858009, + 75.83999633789062, + 50.11200159788132, + 74.68800246715546, + 49.40799996256828, + 76.03199779987335, + 48.576001077890396, + 75.6160020828247, + 50.20799860358238, + 75.45600086450577, + 47.16800153255463, + 76.4480009675026, + 49.27999898791313, + 75.55200159549713, + 51.35999992489815, + 78.75200361013412, + 53.0879981815815, + 75.42400062084198, + 55.48800155520439, + 78.94399762153625, + 52.73599922657013, + 79.19999957084656, + 48.48000034689903, + 74.78400319814682, + 49.82399940490723, + 79.48800176382065, + 49.695998430252075, + 73.91999661922455, + 49.66399818658829, + 74.17599856853485, + 49.18399825692177, + 74.75200295448303, + 48.70399832725525, + 73.60000163316727, + 48.48000034689903, + 77.504001557827, + 51.19999870657921, + 73.66400212049484, + 50.016000866889954, + 74.27199929952621, + 49.34399947524071, + 80.86399734020233, + 50.11200159788132, + 74.01599735021591, + 48.06400090456009, + 78.36800068616867, + 56.2559999525547, + 74.68800246715546, + 49.56800118088722, + 75.13599842786789, + 48.73599857091904, + 74.20799881219864, + 50.84799975156784, + 84.25600081682205, + 48.895999789237976, + 76.19199901819229, + 47.45600000023842, + 82.87999778985977, + 48.895999789237976, + 79.3600007891655, + 48.96000027656555, + 77.82399654388428, + 48.767998814582825, + 75.07199794054031, + 48.86399954557419, + 74.68800246715546, + 49.855999648571014, + 74.14399832487106, + 50.144001841545105, + 78.23999971151352, + 49.215998500585556, + 82.43200182914734, + 50.144001841545105, + 75.45600086450577, + 49.12000149488449, + 74.30399954319, + 49.12000149488449, + 74.5920017361641, + 49.0880012512207, + 74.97599720954895, + 48.8319993019104, + 73.56800138950348, + 49.855999648571014, + 73.56800138950348, + 50.016000866889954, + 82.59200304746628, + 49.50400069355965, + 73.2479989528656, + 54.07999828457832, + 84.25600081682205, + 48.73599857091904, + 75.29599964618683, + 50.97600072622299, + 74.68800246715546, + 50.944000482559204, + 75.29599964618683, + 47.968000173568726, + 76.28799974918365, + 48.48000034689903, + 75.87199658155441, + 50.144001841545105, + 80.51200211048126, + 51.16799846291542, + 73.37599992752075, + 49.34399947524071, + 74.78400319814682, + 48.576001077890396, + 74.97599720954895, + 48.22399839758873, + 73.72800260782242, + 47.90399968624115, + 81.15199953317642, + 49.95200037956238, + 74.52800124883652, + 49.215998500585556, + 74.01599735021591, + 50.144001841545105, + 74.97599720954895, + 48.8319993019104, + 73.88799637556076, + 59.328000992536545, + 75.32799988985062, + 49.855999648571014, + 77.88799703121185, + 49.12000149488449, + 72.73600250482559, + 49.56800118088722, + 76.51200145483017, + 49.056001007556915, + 73.53600114583969, + 48.64000156521797, + 73.56800138950348, + 48.00000041723251, + 75.39200037717819, + 48.895999789237976, + 75.6480023264885, + 49.50400069355965, + 76.73600316047668, + 49.34399947524071, + 76.54400169849396, + 50.75199902057648, + 75.52000135183334, + 50.23999884724617, + 78.5600021481514, + 53.408000618219376, + 74.75200295448303, + 49.984000623226166, + 75.39200037717819, + 50.49600079655647, + 75.3600001335144, + 48.38399961590767, + 76.12799853086472, + 49.375999718904495, + 76.57600194215775, + 48.19199815392494, + 78.27199995517731, + 50.23999884724617, + 76.31999999284744, + 49.82399940490723, + 75.87199658155441, + 50.36799982190132, + 74.5600014925003, + 47.968000173568726, + 77.15199887752533, + 56.2559999525547, + 74.17599856853485, + 47.648001462221146, + 74.91199672222137, + 48.8319993019104, + 76.86399668455124, + 55.16799911856651, + 75.6480023264885, + 49.44000020623207, + 73.11999797821045, + 48.31999912858009, + 75.00799745321274, + 59.167999774217606, + 74.49600100517273, + 47.90399968624115, + 77.18399912118912, + 50.592001527547836, + 75.1039981842041, + 48.19199815392494, + 73.31199944019318, + 49.47200044989586, + 74.23999905586243, + 49.75999891757965, + 74.5920017361641, + 48.128001391887665, + 74.68800246715546, + 48.86399954557419, + 73.95199686288834, + 49.92000013589859, + 82.2720006108284 + ], + "provider": [ + 76.19199901819229, + 106.27199709415436, + 67.74400174617767, + 105.24799674749374, + 66.65600091218948, + 104.70400005578995, + 66.91200286149979, + 105.59999942779541, + 77.18399912118912, + 105.31199723482132, + 69.5360004901886, + 103.71199995279312, + 66.880002617836, + 103.45599800348282, + 66.52799993753433, + 105.27999699115753, + 68.1919977068901, + 166.24000668525696, + 67.77600198984146, + 105.27999699115753, + 67.1359971165657, + 106.97600245475769, + 67.29599833488464, + 104.44799810647964, + 66.84800237417221, + 108.47999900579453, + 69.40799951553345, + 110.68800091743469, + 68.70400160551071, + 104.8320010304451, + 68.7360018491745, + 104.63999956846237, + 67.9360032081604, + 104.8320010304451, + 67.00800359249115, + 106.01600259542465, + 68.06399673223495, + 107.10400342941284, + 66.81600213050842, + 104.19200360774994, + 68.15999746322632, + 108.12799632549286, + 68.06399673223495, + 105.43999820947647, + 67.64800101518631, + 103.58399897813797, + 67.4239993095398, + 108.12799632549286, + 67.87200272083282, + 104.47999835014343, + 67.58400052785873, + 109.18399691581726, + 67.74400174617767, + 111.51999980211258, + 68.06399673223495, + 105.59999942779541, + 67.00800359249115, + 104.19200360774994, + 68.2239979505539, + 110.52799969911575, + 67.10399687290192, + 106.39999806880951, + 67.84000247716904, + 104.60799932479858, + 67.55200028419495, + 106.30399733781815, + 66.91200286149979, + 107.90400207042694, + 73.91999661922455, + 104.00000214576721, + 66.65600091218948, + 104.38399761915207, + 75.13599842786789, + 104.89600151777267, + 66.20799750089645, + 106.11200332641602, + 68.1919977068901, + 107.13600367307663, + 68.70400160551071, + 105.6319996714592, + 67.45599955320358, + 114.30399864912033, + 67.58400052785873, + 105.34399747848511, + 66.68800115585327, + 105.43999820947647, + 67.35999882221222, + 104.41599786281586, + 66.78400188684464, + 104.60799932479858, + 67.52000004053116, + 105.15200346708298, + 68.76800209283829, + 105.95200210809708, + 67.9360032081604, + 103.71199995279312, + 67.61600077152252, + 105.47199845314026, + 68.80000233650208, + 105.76000064611435, + 67.32799857854843, + 105.69600015878677, + 67.19999760389328, + 104.76800054311752, + 67.96800345182419, + 104.47999835014343, + 65.92000275850296, + 107.19999670982361, + 67.03999638557434, + 103.39199751615524, + 66.0799965262413, + 104.89600151777267, + 69.11999732255936, + 105.40799796581268, + 66.49599969387054, + 105.27999699115753, + 67.52000004053116, + 107.71200060844421, + 67.87200272083282, + 105.24799674749374, + 67.23199784755707, + 105.12000322341919, + 66.14399701356888, + 105.53599894046783, + 66.39999896287918, + 105.15200346708298, + 67.26399809122086, + 105.31199723482132, + 66.3679987192154, + 109.27999764680862, + 68.83200258016586, + 105.0880029797554, + 67.4239993095398, + 107.19999670982361, + 67.03999638557434, + 104.38399761915207, + 67.6800012588501, + 104.35199737548828, + 67.52000004053116, + 105.76000064611435, + 67.9360032081604, + 111.39199882745743, + 67.52000004053116, + 106.75200074911118, + 66.880002617836, + 107.96800255775452, + 67.03999638557434, + 108.06400328874588, + 68.12799721956253, + 106.59199953079224, + 73.82400333881378, + 103.29599678516388, + 66.65600091218948, + 104.41599786281586, + 66.59200042486191, + 107.80800133943558, + 68.57600063085556, + 105.18400371074677, + 66.97600334882736, + 104.99200224876404, + 66.97600334882736, + 111.96800321340561, + 67.391999065876, + 104.2879968881607, + 67.16799736022949, + 104.2879968881607, + 66.49599969387054, + 106.27199709415436, + 66.27199798822403, + 108.06400328874588, + 66.11199676990509, + 105.47199845314026, + 68.35199892520905, + 106.175996363163, + 67.391999065876, + 106.75200074911118, + 67.74400174617767, + 106.62399977445602, + 66.23999774456024, + 105.72800040245056, + 66.23999774456024, + 106.65600001811981, + 67.19999760389328, + 104.73600029945374, + 66.0799965262413, + 105.21599650382996, + 67.4239993095398, + 104.35199737548828, + 66.56000018119812, + 104.19200360774994, + 67.35999882221222, + 103.5199984908104, + 67.4239993095398, + 103.93600165843964, + 67.26399809122086, + 105.95200210809708, + 67.6800012588501, + 104.38399761915207, + 67.9360032081604, + 106.97600245475769 + ], + "candidate_backward": [ + 1080.064058303833, + 902.783989906311, + 1086.6880416870117, + 1149.6959924697876, + 774.2720246315002, + 917.6639914512634, + 1132.159948348999, + 1119.5199489593506, + 1024.8960256576538, + 1114.527940750122, + 835.3599905967712, + 1102.944016456604, + 1060.2240562438965, + 1100.8319854736328, + 1114.848017692566, + 1132.7359676361084, + 1142.8159475326538, + 1839.2319679260254, + 1149.0559577941895, + 1130.8480501174927, + 1079.8399448394775, + 1141.2160396575928, + 806.5599799156189, + 1136.8000507354736, + 1147.9359865188599, + 1123.0080127716064, + 810.6880187988281, + 1142.8159475326538, + 827.5200128555298, + 886.6879940032959, + 1107.6480150222778, + 1201.7920017242432, + 1153.7599563598633, + 1083.9999914169312, + 1146.880030632019, + 1216.9599533081055, + 1157.312035560608, + 901.5359878540039, + 1120.5120086669922, + 1098.7520217895508, + 995.4239726066589, + 1135.7439756393433, + 759.4239711761475, + 1126.4959573745728, + 1116.0000562667847, + 1131.7440271377563, + 1130.8799982070923, + 1097.3119735717773, + 790.1120185852051, + 1116.0320043563843, + 1119.6800470352173, + 988.4480237960815, + 1207.6799869537354, + 1141.4400339126587, + 1122.4000453948975, + 1170.240044593811, + 1132.7040195465088, + 1065.2799606323242, + 1126.8160343170166, + 1139.7759914398193, + 779.2959809303284, + 1152.351975440979, + 842.0159816741943, + 1109.3440055847168, + 788.5119915008545, + 1085.4079723358154, + 994.2399859428406, + 1152.448058128357, + 1116.927981376648, + 1258.6239576339722, + 1122.4960088729858, + 1087.3600244522095, + 1112.671971321106, + 1132.6080560684204, + 1133.8880062103271, + 1139.3280029296875, + 802.1759986877441, + 1138.7840509414673, + 782.5279831886292, + 1120.09596824646, + 1100.2880334854126, + 999.072015285492, + 1159.6800088882446, + 1123.039960861206, + 1152.448058128357, + 1117.6320314407349, + 1128.640055656433, + 1087.231993675232, + 1143.5199975967407, + 1159.1360569000244, + 1119.2320585250854, + 1140.3520107269287, + 1116.2559986114502, + 1111.232042312622, + 1138.3039951324463, + 1192.5439834594727, + 1113.055944442749, + 1114.8159503936768, + 1101.8879413604736, + 1136.896014213562, + 849.407970905304, + 1113.4400367736816, + 1113.0880117416382, + 1121.8880414962769, + 1035.8400344848633, + 1122.431993484497, + 1111.2639904022217, + 1170.8159446716309, + 1131.168007850647, + 1126.4640092849731, + 1137.3440027236938, + 1152.9279947280884, + 1121.791958808899, + 1150.3039598464966, + 1139.0719413757324, + 1139.8080587387085, + 1150.5600214004517, + 1289.6959781646729, + 807.0719838142395, + 1062.7520084381104, + 1135.6799602508545, + 1119.4239854812622, + 1145.4720497131348, + 1074.8800039291382, + 1220.9279537200928, + 1165.727972984314, + 1163.167953491211, + 1336.7040157318115, + 1072.3520517349243, + 940.7680034637451, + 1067.4879550933838, + 1139.9999856948853, + 1042.6239967346191, + 1128.159999847412, + 1100.0319719314575, + 1123.0720281600952, + 1197.4719762802124, + 1144.8639631271362, + 1098.464012145996, + 1081.5999507904053, + 1126.4640092849731, + 1138.11194896698, + 1194.3360567092896, + 1162.7839803695679, + 803.2960295677185, + 1108.672022819519, + 760.0640058517456, + 1120.5120086669922, + 1000.4160404205322, + 1108.415961265564, + 1085.8880281448364, + 1133.1839561462402, + 1132.6719522476196, + 1132.383942604065, + 1098.3680486679077, + 1129.696011543274, + 1076.032042503357, + 905.4080247879028, + 1116.5119409561157, + 1131.5200328826904, + 1123.136043548584, + 1128.3520460128784, + 806.1439990997314, + 894.4640159606934, + 1127.7120113372803, + 964.7359848022461, + 1118.1119680404663, + 1082.0480585098267, + 1152.8639793395996, + 1136.896014213562, + 825.0880241394043, + 871.936023235321, + 993.5680031776428, + 1131.4239501953125, + 1185.5679750442505, + 1136.2240314483643, + 776.1600017547607, + 1183.6479902267456, + 996.6719746589661, + 1074.8159885406494, + 1182.528018951416, + 1100.000023841858, + 1177.9839992523193, + 946.2720155715942, + 1169.119954109192, + 1127.679944038391, + 776.6079902648926, + 1139.0080451965332, + 1132.64000415802, + 1119.488000869751, + 1100.3520488739014, + 998.1439709663391, + 1150.1439809799194, + 1100.864052772522, + 1097.599983215332, + 1091.0719633102417, + 1131.2960386276245, + 1135.3919506072998, + 1259.9040269851685, + 1109.663963317871 + ], + "provider_backward": [ + 615.552008152008, + 584.9599838256836, + 606.2719821929932, + 566.5280222892761, + 589.6959900856018, + 562.175989151001, + 649.7600078582764, + 595.9360003471375, + 632.4800252914429, + 597.7920293807983, + 569.599986076355, + 563.4559988975525, + 565.4399991035461, + 574.8479962348938, + 570.0479745864868, + 587.8720283508301, + 595.52001953125, + 546.8479990959167, + 621.6639876365662, + 627.9039978981018, + 571.4880228042603, + 578.6240100860596, + 585.3440165519714, + 599.1359949111938, + 436.2879991531372, + 580.1920294761658, + 554.3040037155151, + 568.0959820747375, + 440.5759871006012, + 434.2080056667328, + 607.3600053787231, + 535.5839729309082, + 520.1280117034912, + 638.5279893875122, + 464.86398577690125, + 582.3360085487366, + 588.2880091667175, + 583.1360220909119, + 594.7520136833191, + 579.8720121383667, + 573.2160210609436, + 582.3040008544922, + 411.45598888397217, + 597.0879793167114, + 644.320011138916, + 562.720000743866, + 587.6160264015198, + 596.0639715194702, + 596.3199734687805, + 553.2799959182739, + 604.095995426178, + 585.1839780807495, + 582.5600028038025, + 416.9920086860657, + 575.1680135726929, + 603.6800146102905, + 565.2480125427246, + 560.5760216712952, + 604.6720147132874, + 587.5200033187866, + 564.5760297775269, + 575.3920078277588, + 546.3039875030518, + 607.1360111236572, + 425.2159893512726, + 589.024007320404, + 598.7840294837952, + 566.1439895629883, + 623.8399744033813, + 561.568021774292, + 574.783980846405, + 662.015974521637, + 583.1040143966675, + 560.1279735565186, + 573.4080076217651, + 566.4640069007874, + 570.7200169563293, + 564.0320181846619, + 594.6879982948303, + 588.2880091667175, + 545.0559854507446, + 590.5920267105103, + 410.3679955005646, + 560.1279735565186, + 575.9680271148682, + 615.9359812736511, + 596.127986907959, + 564.736008644104, + 422.2399890422821, + 594.6559906005859, + 576.6720175743103, + 588.6399745941162, + 553.9199709892273, + 589.6639823913574, + 582.431972026825, + 690.2719736099243, + 556.1919808387756, + 545.7919836044312, + 567.1679973602295, + 573.6640095710754, + 403.1040072441101, + 565.7920241355896, + 483.6159944534302, + 576.3199925422668, + 562.4960064888, + 548.0639934539795, + 413.05598616600037, + 559.1359734535217, + 594.8479771614075, + 452.57601141929626, + 646.399974822998, + 515.6159996986389, + 588.9599919319153, + 554.5920133590698, + 574.5279788970947, + 441.9200122356415, + 580.0639986991882, + 585.8880281448364, + 425.1199960708618, + 560.9279870986938, + 575.4240155220032, + 586.3040089607239, + 589.6639823913574, + 644.0960168838501, + 567.9360032081604, + 709.1839909553528, + 645.3120112419128, + 587.9359841346741, + 581.3760161399841, + 589.9839997291565, + 565.280020236969, + 558.5600137710571, + 621.5999722480774, + 808.3840012550354, + 423.99999499320984, + 534.4640016555786, + 575.5199790000916, + 583.296000957489, + 574.8479962348938, + 398.49600195884705, + 615.6799793243408, + 558.2720041275024, + 360.83200573921204, + 592.3200249671936, + 586.0159993171692, + 577.0879983901978, + 580.7679891586304, + 416.54399037361145, + 576.7359733581543, + 645.6639766693115, + 633.1520080566406, + 570.0799822807312, + 631.712019443512, + 540.8639907836914, + 551.8400073051453, + 594.976007938385, + 587.6479744911194, + 581.0880064964294, + 586.0159993171692, + 553.9199709892273, + 508.54402780532837, + 617.8879737854004, + 546.720027923584, + 579.9040198326111, + 590.2079939842224, + 559.1679811477661, + 584.5760107040405, + 576.1280059814453, + 570.1119899749756, + 585.5680108070374, + 555.1360249519348, + 582.9759836196899, + 559.6160292625427, + 579.2639851570129, + 553.2479882240295, + 559.3600273132324, + 358.43199491500854, + 536.6399884223938, + 576.2559771537781, + 458.14400911331177, + 572.4800229072571, + 583.9359760284424, + 584.2559933662415, + 575.7439732551575, + 596.0000157356262, + 566.5599703788757, + 645.5360054969788, + 574.2719769477844, + 568.4159994125366, + 569.3439841270447, + 423.5199987888336, + 552.511990070343, + 407.45601058006287, + 587.3919725418091, + 566.2720203399658, + 598.5280275344849, + 657.9520106315613, + 600.1279950141907, + 551.8720149993896, + 606.112003326416 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 66.03200174868107, + "candidate_gbps": 2001.3452341332304, + "candidate_peak_mib": 42.01025390625, + "provider_us": 90.2399979531765, + "provider_gbps": 1464.459607684955, + "provider_peak_mib": 84.0205078125, + "candidate_backward_us": 1120.303988456726, + "candidate_backward_gbps": 117.9615830717938, + "candidate_backward_peak_mib": 84.66650390625, + "provider_backward_us": 576.0480165481567, + "provider_backward_gbps": 229.4128756694577, + "provider_backward_peak_mib": 168.041015625 + }, + { + "op": "adaln_gate_residual", + "case": "S=32768", + "backend": "H3GateResidualCudaOp", + "bytes": 1056964608, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 238.39999735355377, + 256.99201226234436, + 231.455996632576, + 257.31199979782104, + 230.75200617313385, + 257.0880055427551, + 232.4800044298172, + 257.24801421165466, + 231.455996632576, + 258.30399990081787, + 234.17599499225616, + 255.90398907661438, + 232.41600394248962, + 257.31199979782104, + 232.57599771022797, + 255.71200251579285, + 230.71999847888947, + 256.79999589920044, + 232.63999819755554, + 264.2880082130432, + 231.29600286483765, + 255.87201118469238, + 232.16000199317932, + 257.1519911289215, + 232.4800044298172, + 258.0159902572632, + 230.81600666046143, + 257.2160065174103, + 232.35200345516205, + 256.1599910259247, + 231.77599906921387, + 255.87201118469238, + 232.89600014686584, + 257.60000944137573, + 232.12799429893494, + 255.96800446510315, + 230.3999960422516, + 257.56800174713135, + 230.6559979915619, + 257.1839988231659, + 240.03200232982635, + 257.31199979782104, + 231.48800432682037, + 257.4079930782318, + 230.6240051984787, + 256.03199005126953, + 230.43200373649597, + 256.25601410865784, + 231.36000335216522, + 257.4400007724762, + 230.71999847888947, + 258.0159902572632, + 232.41600394248962, + 256.9279968738556, + 233.024001121521, + 255.840003490448, + 231.455996632576, + 257.1200132369995, + 232.7679991722107, + 258.2080066204071, + 232.09600150585175, + 257.3759853839874, + 234.17599499225616, + 256.6719949245453, + 232.83199965953827, + 256.9279968738556, + 231.04000091552734, + 257.5039863586426, + 232.86400735378265, + 258.04799795150757, + 232.54400491714478, + 257.53599405288696, + 232.31999576091766, + 258.5279941558838, + 233.15200209617615, + 258.4640085697174, + 232.41600394248962, + 257.1200132369995, + 231.80800676345825, + 259.13599133491516, + 232.1919947862625, + 256.3839852809906, + 233.11999440193176, + 257.88798928260803, + 237.59999871253967, + 256.79999589920044, + 232.12799429893494, + 256.99201226234436, + 231.99999332427979, + 257.4079930782318, + 231.83999955654144, + 237.7600073814392, + 235.07200181484222, + 256.1280131340027, + 231.3919961452484, + 258.08000564575195, + 232.5119972229004, + 257.1200132369995, + 232.44799673557281, + 258.11201333999634, + 232.06399381160736, + 236.67199909687042, + 231.80800676345825, + 256.9600045681, + 231.3919961452484, + 257.53599405288696, + 231.6800057888031, + 257.1519911289215, + 230.9119999408722, + 255.51998615264893, + 232.70399868488312, + 265.53601026535034, + 230.17600178718567, + 255.10400533676147, + 230.68800568580627, + 256.48000836372375, + 239.9040013551712, + 256.3839852809906, + 230.6240051984787, + 256.44800066947937, + 230.97600042819977, + 255.90398907661438, + 232.12799429893494, + 257.4400007724762, + 230.49600422382355, + 258.4959864616394, + 231.4240038394928, + 256.73601031303406, + 232.1919947862625, + 256.28799200057983, + 231.1680018901825, + 255.96800446510315, + 231.26399517059326, + 258.1759989261627, + 232.86400735378265, + 256.9279968738556, + 231.99999332427979, + 257.9840123653412, + 231.29600286483765, + 256.0639977455139, + 231.7119985818863, + 257.1200132369995, + 231.77599906921387, + 256.54399394989014, + 235.6799989938736, + 258.432000875473, + 232.70399868488312, + 264.51200246810913, + 232.80000686645508, + 257.79199600219727, + 234.0800017118454, + 256.7040026187897, + 231.61600530147552, + 258.59200954437256, + 231.48800432682037, + 258.33600759506226, + 232.89600014686584, + 260.80000400543213, + 231.9680005311966, + 257.1200132369995, + 232.70399868488312, + 256.44800066947937, + 231.55200481414795, + 255.5519938468933, + 231.1359941959381, + 258.08000564575195, + 233.11999440193176, + 257.79199600219727, + 232.2240024805069, + 265.79201221466064, + 232.28800296783447, + 258.65599513053894, + 233.2800030708313, + 255.16799092292786, + 232.09600150585175, + 258.2719922065735, + 231.10400140285492, + 257.2160065174103, + 234.17599499225616, + 257.9199969768524, + 236.4799976348877, + 258.30399990081787, + 232.35200345516205, + 257.9840123653412, + 231.1999946832657, + 256.76798820495605, + 231.64799809455872, + 256.8959891796112, + 232.38399624824524, + 258.62398743629456, + 231.55200481414795, + 256.19199872016907, + 231.55200481414795, + 257.9520046710968, + 232.67200589179993, + 256.9279968738556, + 232.2240024805069, + 257.2160065174103, + 232.92799293994904, + 258.62398743629456 + ], + "provider": [ + 416.3840115070343, + 448.35200905799866, + 415.6480133533478, + 447.4239945411682, + 416.703999042511, + 449.66399669647217, + 418.3039963245392, + 449.7919976711273, + 416.4159893989563, + 446.75201177597046, + 417.4399971961975, + 449.6000111103058, + 417.7280068397522, + 448.2559859752655, + 417.34400391578674, + 448.15999269485474, + 413.2159948348999, + 449.5680034160614, + 415.74400663375854, + 449.0559995174408, + 413.6640131473541, + 450.27199387550354, + 415.23200273513794, + 448.7360119819641, + 416.159987449646, + 451.03999972343445, + 416.3520038127899, + 449.8879909515381, + 415.9359931945801, + 448.7360119819641, + 415.74400663375854, + 448.2879936695099, + 415.45599699020386, + 448.7360119819641, + 417.08800196647644, + 450.3679871559143, + 413.12000155448914, + 447.7759897708893, + 418.2080030441284, + 449.24798607826233, + 419.96800899505615, + 449.24798607826233, + 415.5519902706146, + 447.32800126075745, + 416.3520038127899, + 446.27198576927185, + 419.51999068260193, + 449.69600439071655, + 415.8079922199249, + 448.92799854278564, + 417.279988527298, + 448.2240080833435, + 414.4960045814514, + 451.26399397850037, + 415.74400663375854, + 449.7919976711273, + 415.5200123786926, + 455.7119905948639, + 418.720006942749, + 451.200008392334, + 416.03198647499084, + 445.82399725914, + 417.4399971961975, + 450.01599192619324, + 417.4399971961975, + 446.8800127506256, + 416.128009557724, + 450.49598813056946, + 413.2159948348999, + 448.15999269485474, + 413.7600064277649, + 448.7679898738861, + 418.2400107383728, + 447.61601090431213, + 416.1919951438904, + 451.35998725891113, + 413.1839871406555, + 451.9680142402649, + 415.2640104293823, + 452.5440037250519, + 418.33600401878357, + 449.3120014667511, + 419.1359877586365, + 447.9680061340332, + 413.31198811531067, + 450.23998618125916, + 418.08000206947327, + 448.4800100326538, + 417.4720048904419, + 448.86401295661926, + 416.9279932975769, + 425.50399899482727, + 414.40001130104065, + 452.1600008010864, + 414.68799114227295, + 446.5920031070709, + 417.279988527298, + 449.535995721817, + 415.2640104293823, + 448.67199659347534, + 418.43199729919434, + 423.99999499320984, + 416.79999232292175, + 448.7040042877197, + 415.583997964859, + 446.399986743927, + 415.5519902706146, + 448.7679898738861, + 415.8399999141693, + 460.31999588012695, + 416.0960018634796, + 447.32800126075745, + 415.3600037097931, + 449.44000244140625, + 416.0960018634796, + 446.55999541282654, + 413.536012172699, + 447.07199931144714, + 417.63201355934143, + 447.3919868469238, + 416.4479970932007, + 446.9119906425476, + 417.5040125846863, + 446.8800127506256, + 416.0960018634796, + 446.9119906425476, + 417.82400012016296, + 449.3759870529175, + 418.11200976371765, + 451.1680006980896, + 415.9359931945801, + 451.200008392334, + 417.56799817085266, + 450.3999948501587, + 417.91999340057373, + 447.1679925918579, + 414.65601325035095, + 449.44000244140625, + 415.2640104293823, + 452.06400752067566, + 416.9600009918213, + 448.2240080833435, + 415.039986371994, + 449.8240053653717, + 418.2719886302948, + 450.080007314682, + 415.16798734664917, + 448.2559859752655, + 417.6639914512634, + 449.6319890022278, + 419.74401473999023, + 450.9119987487793, + 412.8960072994232, + 448.7040042877197, + 419.16799545288086, + 452.4480104446411, + 415.71199893951416, + 446.8480050563812, + 415.1040017604828, + 452.0319998264313, + 415.23200273513794, + 451.58401131629944, + 417.5040125846863, + 457.0879936218262, + 419.840008020401, + 452.2559940814972, + 413.85599970817566, + 450.20800828933716, + 418.8160002231598, + 449.69600439071655, + 416.28798842430115, + 450.23998618125916, + 417.60000586509705, + 449.72801208496094, + 412.57598996162415, + 447.55199551582336, + 415.77601432800293, + 447.2320079803467, + 414.91198539733887, + 452.2559940814972, + 416.31999611854553, + 451.58401131629944, + 416.51201248168945, + 447.6799964904785, + 417.34400391578674, + 458.0160081386566, + 414.94399309158325, + 447.7440118789673, + 416.1919951438904, + 450.01599192619324, + 417.60000586509705, + 451.6479969024658, + 414.94399309158325, + 448.06399941444397, + 416.76801443099976, + 450.1439929008484, + 412.86399960517883, + 449.44000244140625, + 415.1040017604828, + 447.55199551582336 + ], + "candidate_backward": [ + 1093.567967414856, + 1175.4560470581055, + 1099.8400449752808, + 1168.3520078659058, + 1133.5680484771729, + 1120.255947113037, + 1115.1360273361206, + 1148.0640172958374, + 1070.1119899749756, + 1165.1840209960938, + 1102.8800010681152, + 1165.3120517730713, + 1087.5200033187866, + 1191.9679641723633, + 1131.55198097229, + 1146.8479633331299, + 1159.1039896011353, + 1120.2239990234375, + 1096.7040061950684, + 1147.9040384292603, + 1107.9360246658325, + 1166.3039922714233, + 1088.5440111160278, + 1160.159945487976, + 1136.672019958496, + 1174.5280027389526, + 1123.3279705047607, + 1174.5599508285522, + 1101.8240451812744, + 1128.4480094909668, + 1116.1279678344727, + 1123.3279705047607, + 965.6320214271545, + 1172.4799871444702, + 1120.1599836349487, + 1109.9519729614258, + 1070.080041885376, + 1115.1679754257202, + 1093.664050102234, + 1156.1599969863892, + 1113.0880117416382, + 1180.7680130004883, + 1117.2480583190918, + 1134.592056274414, + 1100.8319854736328, + 1168.3520078659058, + 1121.3120222091675, + 1161.1839532852173, + 1101.8240451812744, + 1092.6079750061035, + 1124.384045600891, + 1121.3760375976562, + 1062.9440546035767, + 1117.1519756317139, + 1113.0880117416382, + 1154.0800333023071, + 1093.6319828033447, + 957.1520090103149, + 1134.6559524536133, + 1124.384045600891, + 1113.0880117416382, + 1114.2079830169678, + 1144.4799900054932, + 1221.6320037841797, + 1070.1440572738647, + 1022.9760408401489, + 1094.655990600586, + 1108.0000400543213, + 1144.8320150375366, + 1115.1360273361206, + 1071.1040496826172, + 1182.6879978179932, + 1089.568018913269, + 984.063982963562, + 966.6560292243958, + 1121.3120222091675, + 1132.5440406799316, + 1154.8800468444824, + 1139.7119760513306, + 1165.3120517730713, + 928.76797914505, + 1102.8480529785156, + 1099.776029586792, + 1105.9199571609497, + 1091.5839672088623, + 1160.1920127868652, + 1082.368016242981, + 1011.7119550704956, + 1104.8959493637085, + 1154.0800333023071, + 1105.9199571609497, + 987.1360063552856, + 1062.9440546035767, + 1165.3120517730713, + 1104.7680377960205, + 1121.3120222091675, + 1089.568018913269, + 1159.2639684677124, + 1063.9359951019287, + 1153.0239582061768, + 1122.3039627075195, + 1086.4959955215454, + 1104.9280166625977, + 1164.28804397583, + 1111.0399961471558, + 1090.559959411621, + 1093.6319828033447, + 1171.455979347229, + 1100.8000373840332, + 1139.7440433502197, + 1124.2879629135132, + 1205.183982849121, + 1134.6240043640137, + 1162.2400283813477, + 1090.656042098999, + 1153.9520025253296, + 1127.3599863052368, + 1162.2719764709473, + 1115.2000427246094, + 1120.2880144119263, + 1047.551989555359, + 1216.5119647979736, + 1174.496054649353, + 1144.8639631271362, + 1118.2399988174438, + 1222.656011581421, + 1079.2640447616577, + 1142.848014831543, + 1101.8240451812744, + 1122.3360300064087, + 1078.2400369644165, + 1136.639952659607, + 969.7279930114746, + 1098.7520217895508, + 1146.880030632019, + 1162.2719764709473, + 1083.4239721298218, + 1112.064003944397, + 1090.432047843933, + 1154.2079448699951, + 1135.6159448623657, + 1135.6159448623657, + 1137.6320123672485, + 1117.1519756317139, + 1095.7119464874268, + 1118.2080507278442, + 1117.184042930603, + 1098.7839698791504, + 1107.9360246658325, + 1111.0719442367554, + 1119.1359758377075, + 1103.8399934768677, + 1114.1120195388794, + 1139.7440433502197, + 1116.1919832229614, + 1128.4159421920776, + 1074.1440057754517, + 1106.8799495697021, + 1107.9360246658325, + 1117.184042930603, + 1104.8959493637085, + 1110.0159883499146, + 1138.6879682540894, + 1112.064003944397, + 1062.9119873046875, + 1114.1120195388794, + 1125.3759860992432, + 1214.4960165023804, + 1122.3039627075195, + 1126.3999938964844, + 1109.0240478515625, + 1103.9040088653564, + 1115.1039600372314, + 1182.752013206482, + 1122.3039627075195, + 1153.9839506149292, + 1089.5040035247803, + 983.0080270767212, + 1122.3039627075195, + 1073.15194606781, + 1101.7919778823853, + 1133.5680484771729, + 1076.192021369934, + 1273.8560438156128, + 1092.6079750061035, + 1111.0080480575562, + 1069.0560340881348, + 1244.1600561141968, + 1093.6319828033447, + 1043.4240102767944, + 1116.1919832229614, + 1002.4960041046143, + 1081.3759565353394, + 1253.3440589904785, + 1118.175983428955, + 1141.759991645813, + 1123.36003780365, + 1099.8400449752808, + 989.0879988670349, + 1076.2239694595337 + ], + "provider_backward": [ + 2244.54402923584, + 2241.5359020233154, + 2233.3760261535645, + 2236.448049545288, + 2243.6161041259766, + 2244.640111923218, + 2245.6319332122803, + 2235.2960109710693, + 2242.5599098205566, + 2239.5200729370117, + 2249.376058578491, + 2238.3999824523926, + 2241.5359020233154, + 2240.544080734253, + 2242.4960136413574, + 2236.3200187683105, + 2238.6560440063477, + 2235.424041748047, + 2237.4720573425293, + 2231.2960624694824, + 2305.0239086151123, + 2241.503953933716, + 2240.511894226074, + 2229.248046875, + 2239.583969116211, + 2243.583917617798, + 2238.464117050171, + 2235.424041748047, + 2243.583917617798, + 2235.424041748047, + 2240.511894226074, + 2233.344078063965, + 2242.527961730957, + 2243.4239387512207, + 2243.6161041259766, + 2240.544080734253, + 2241.568088531494, + 2231.3599586486816, + 2245.66388130188, + 2235.4559898376465, + 2229.2799949645996, + 2237.4401092529297, + 2244.607925415039, + 2235.3599071502686, + 2238.4960651397705, + 2239.6481037139893, + 2239.487886428833, + 2236.3839149475098, + 2241.6000366210938, + 2235.424041748047, + 2234.272003173828, + 2236.3839149475098, + 2240.511894226074, + 2237.407922744751, + 2247.551918029785, + 2311.1679553985596, + 2240.511894226074, + 2235.3599071502686, + 2238.4960651397705, + 2237.4720573425293, + 2242.527961730957, + 2237.3759746551514, + 2243.5519695281982, + 2237.5359535217285, + 2237.4401092529297, + 2236.4161014556885, + 2237.4401092529297, + 2233.5360050201416, + 2235.424041748047, + 2238.3999824523926, + 2240.384101867676, + 2240.4799461364746, + 2244.256019592285, + 2238.464117050171, + 2235.3920936584473, + 2234.3039512634277, + 2244.607925415039, + 2231.328010559082, + 2239.487886428833, + 2238.431930541992, + 2237.312078475952, + 2234.4000339508057, + 2247.7118968963623, + 2237.4401092529297, + 2246.6559410095215, + 2241.5359020233154, + 2243.583917617798, + 2245.66388130188, + 2244.640111923218, + 2234.368085861206, + 2242.464065551758, + 2243.6161041259766, + 2235.424041748047, + 2235.3920936584473, + 2246.6559410095215, + 2235.3599071502686, + 2241.568088531494, + 2237.567901611328, + 2243.6161041259766, + 2237.4720573425293, + 2245.6319332122803, + 2235.3599071502686, + 2241.6000366210938, + 2234.368085861206, + 2240.511894226074, + 2239.4559383392334, + 2247.7760314941406, + 2243.5200214385986, + 2239.5200729370117, + 2236.4799976348877, + 2239.5520210266113, + 2235.3920936584473, + 2240.544080734253, + 2243.5519695281982, + 2239.6159172058105, + 2241.5359020233154, + 2243.583917617798, + 2233.3760261535645, + 2241.6000366210938, + 2235.327959060669, + 2237.4401092529297, + 2241.312026977539, + 2239.3600940704346, + 2317.280054092407, + 2241.5359020233154, + 2239.487886428833, + 2243.6161041259766, + 2235.424041748047, + 2242.5599098205566, + 2263.040065765381, + 2236.448049545288, + 2235.3920936584473, + 2247.7118968963623, + 2232.3520183563232, + 2238.4960651397705, + 2233.344078063965, + 2239.5200729370117, + 2235.23211479187, + 2236.448049545288, + 2233.4399223327637, + 2238.464117050171, + 2233.344078063965, + 2244.352102279663, + 2235.3920936584473, + 2241.5359020233154, + 2239.5200729370117, + 2237.4401092529297, + 2241.5359020233154, + 2238.464117050171, + 2236.0639572143555, + 2245.66388130188, + 2234.368085861206, + 2231.2960624694824, + 2241.568088531494, + 2235.3599071502686, + 2235.3920936584473, + 2247.6799488067627, + 2240.19193649292, + 2239.5200729370117, + 2237.4720573425293, + 2241.5359020233154, + 2243.6161041259766, + 2241.568088531494, + 2235.2640628814697, + 2242.5599098205566, + 2233.344078063965, + 2245.66388130188, + 2235.3920936584473, + 2245.6319332122803, + 2238.4960651397705, + 2239.487886428833, + 2238.208055496216, + 2238.4960651397705, + 2232.3520183563232, + 2245.568037033081, + 2237.4401092529297, + 2240.544080734253, + 2237.4401092529297, + 2235.424041748047, + 2239.487886428833, + 2239.5200729370117, + 2235.424041748047, + 2244.607925415039, + 2242.2399520874023, + 2239.5200729370117, + 2231.2960624694824, + 2235.424041748047, + 2233.344078063965, + 2240.5760288238525, + 2239.5520210266113, + 2242.5920963287354, + 2235.424041748047, + 2234.4000339508057, + 2235.3920936584473, + 2239.487886428833, + 2238.464117050171, + 2245.6319332122803, + 2232.3200702667236, + 2241.6000366210938, + 2241.5359020233154 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 239.1519993543625, + "candidate_gbps": 4419.635256462344, + "candidate_peak_mib": 336.0, + "provider_us": 421.984001994133, + "provider_gbps": 2504.7504241990086, + "provider_peak_mib": 672.0, + "candidate_backward_us": 1116.6719794273376, + "candidate_backward_gbps": 946.5309665440361, + "candidate_backward_peak_mib": 672.64599609375, + "provider_backward_us": 2239.487886428833, + "provider_backward_gbps": 471.96710212417065, + "provider_backward_peak_mib": 1344.0 + }, + { + "op": "adaln_gate_residual", + "case": "S=131072", + "backend": "H3GateResidualCudaOp", + "bytes": 4227858432, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 870.0159788131714, + 888.3200287818909, + 865.119993686676, + 887.9680037498474, + 861.3759875297546, + 888.5120153427124, + 864.3519878387451, + 887.328028678894, + 861.2800240516663, + 889.2160058021545, + 861.7600202560425, + 885.5040073394775, + 860.9279990196228, + 885.7280015945435, + 860.9279990196228, + 890.1119828224182, + 861.3759875297546, + 886.3679766654968, + 860.8959913253784, + 887.2320055961609, + 863.7760281562805, + 887.328028678894, + 860.863983631134, + 886.3679766654968, + 860.9920144081116, + 887.9680037498474, + 862.559974193573, + 891.7440176010132, + 862.9119992256165, + 886.784017086029, + 861.6319894790649, + 973.0240106582642, + 875.9999871253967, + 977.5679707527161, + 860.6079816818237, + 888.1279826164246, + 859.2640161514282, + 888.4479999542236, + 862.5919818878174, + 894.0479755401611, + 862.7200126647949, + 887.9680037498474, + 861.5040183067322, + 889.2480134963989, + 860.6399893760681, + 894.4960236549377, + 864.6079897880554, + 890.7840251922607, + 863.2959723472595, + 886.0800266265869, + 862.1119856834412, + 890.2720212936401, + 863.0080223083496, + 888.9600038528442, + 862.4320030212402, + 888.0640268325806, + 860.0640296936035, + 890.3040289878845, + 860.5440258979797, + 891.7440176010132, + 860.2240085601807, + 888.0640268325806, + 867.6480054855347, + 888.3839845657349, + 861.8239760398865, + 894.4960236549377, + 862.5280261039734, + 888.1279826164246, + 859.7760200500488, + 893.9200043678284, + 863.9360070228577, + 890.3679847717285, + 863.2640242576599, + 889.7920250892639, + 862.6239895820618, + 888.8639807701111, + 862.496018409729, + 889.5999789237976, + 861.5040183067322, + 890.3040289878845, + 860.319972038269, + 891.0719752311707, + 860.9920144081116, + 887.7760171890259, + 861.4720106124878, + 886.8160247802734, + 860.5440258979797, + 890.720009803772, + 860.3839874267578, + 889.6639943122864, + 861.8559837341309, + 889.9840116500854, + 863.2320165634155, + 890.3359770774841, + 862.3999953269958, + 888.8959884643555, + 861.5679740905762, + 888.3200287818909, + 861.1199855804443, + 891.4240002632141, + 860.80002784729, + 888.8959884643555, + 860.2240085601807, + 887.8399729728699, + 861.1840009689331, + 889.3120288848877, + 861.2800240516663, + 889.7920250892639, + 860.4480028152466, + 889.5040154457092, + 862.8479838371277, + 888.1279826164246, + 862.6559972763062, + 889.4079923629761, + 859.7760200500488, + 888.0640268325806, + 859.5200181007385, + 888.5440230369568, + 862.0799779891968, + 887.8719806671143, + 862.6880049705505, + 890.175998210907, + 860.863983631134, + 887.6799941062927, + 861.5040183067322, + 887.55202293396, + 859.391987323761, + 886.9119882583618, + 860.2240085601807, + 890.6880021095276, + 867.7759766578674, + 890.2400135993958, + 861.4720106124878, + 886.6559863090515, + 862.2400164604187, + 888.1919980049133, + 861.9840145111084, + 915.1679873466492, + 864.7680282592773, + 896.992027759552, + 863.4240031242371, + 888.6719942092896, + 861.1199855804443, + 889.4400000572205, + 866.0799860954285, + 901.6640186309814, + 864.2879724502563, + 894.3679928779602, + 861.024022102356, + 894.8479890823364, + 867.4880266189575, + 888.2240056991577, + 860.7360124588013, + 887.615978717804, + 864.6399974822998, + 892.5759792327881, + 864.8959994316101, + 887.3599767684937, + 868.6400055885315, + 889.4720077514648, + 864.9920225143433, + 890.8159732818604, + 861.7920279502869, + 887.9680037498474, + 863.0719780921936, + 887.4880075454712, + 863.0399703979492, + 891.0719752311707, + 861.2160086631775, + 890.496015548706, + 863.647997379303, + 888.7680172920227, + 860.80002784729, + 889.7920250892639, + 870.0159788131714, + 889.4720077514648, + 860.6079816818237, + 887.8719806671143, + 858.9760065078735, + 888.0959749221802, + 859.1679930686951, + 891.8399810791016, + 862.8799915313721, + 888.1279826164246, + 865.8879995346069, + 892.7680253982544, + 872.9280233383179, + 898.2399702072144, + 867.6480054855347, + 890.5280232429504, + 867.4880266189575, + 888.159990310669, + 863.7440204620361, + 888.9279961585999, + 877.5039911270142, + 889.5679712295532, + 860.4480028152466, + 887.1999979019165, + 861.6319894790649, + 888.3519768714905 + ], + "provider": [ + 1599.1679430007935, + 1633.1520080566406, + 1599.1679430007935, + 1628.991961479187, + 1595.4560041427612, + 1627.8719902038574, + 1602.6240587234497, + 1630.2399635314941, + 1592.9280519485474, + 1626.7199516296387, + 1595.7119464874268, + 1627.7120113372803, + 1602.112054824829, + 1631.6800117492676, + 1592.5120115280151, + 1627.616047859192, + 1596.2560176849365, + 1630.303978919983, + 1592.8000211715698, + 1628.4799575805664, + 1596.3200330734253, + 1628.607988357544, + 1595.3279733657837, + 1630.0159692764282, + 1601.8240451812744, + 1627.776026725769, + 1592.8959846496582, + 1630.7519674301147, + 1598.7839698791504, + 1624.1919994354248, + 1596.8639850616455, + 1678.7519454956055, + 1603.4560203552246, + 1663.8720035552979, + 1601.151943206787, + 1626.6239881515503, + 1594.7840213775635, + 1628.5760402679443, + 1595.2320098876953, + 1631.2320232391357, + 1601.7919778823853, + 1633.4400177001953, + 1596.351981163025, + 1628.3520460128784, + 1593.9199924468994, + 1634.112000465393, + 1600.4799604415894, + 1638.8800144195557, + 1593.1199789047241, + 1627.9360055923462, + 1596.2560176849365, + 1631.0399770736694, + 1595.7759618759155, + 1629.5679807662964, + 1594.7200059890747, + 1625.7599592208862, + 1592.9919481277466, + 1629.2799711227417, + 1595.1039791107178, + 1627.616047859192, + 1596.8960523605347, + 1631.4239501953125, + 1598.207950592041, + 1631.6479444503784, + 1594.6240425109863, + 1631.5840482711792, + 1593.34397315979, + 1629.9200057983398, + 1596.0639715194702, + 1633.0560445785522, + 1595.4879522323608, + 1636.288046836853, + 1595.2320098876953, + 1629.6639442443848, + 1597.4719524383545, + 1633.0560445785522, + 1597.5359678268433, + 1627.8079748153687, + 1595.8720445632935, + 1632.7040195465088, + 1597.1200466156006, + 1629.7919750213623, + 1597.2479581832886, + 1627.5839805603027, + 1595.744013786316, + 1631.0720443725586, + 1595.2320098876953, + 1632.383942604065, + 1594.4000482559204, + 1630.7200193405151, + 1598.2400178909302, + 1632.256031036377, + 1597.599983215332, + 1630.4320096969604, + 1597.4719524383545, + 1632.416009902954, + 1594.9759483337402, + 1631.6159963607788, + 1595.263957977295, + 1624.4479417800903, + 1595.52001953125, + 1633.3119869232178, + 1591.3280248641968, + 1631.0399770736694, + 1595.2320098876953, + 1631.168007850647, + 1596.4479446411133, + 1632.3200464248657, + 1597.3440408706665, + 1628.5439729690552, + 1597.2800254821777, + 1630.9759616851807, + 1592.7679538726807, + 1630.911946296692, + 1592.960000038147, + 1625.7599592208862, + 1595.1039791107178, + 1628.0640363693237, + 1594.8480367660522, + 1631.1039924621582, + 1595.6159830093384, + 1628.7039518356323, + 1593.183994293213, + 1632.64000415802, + 1596.1600542068481, + 1630.0480365753174, + 1591.5839672088623, + 1629.8240423202515, + 1596.992015838623, + 1628.5760402679443, + 1596.351981163025, + 1631.168007850647, + 1599.552035331726, + 1629.696011543274, + 1594.7200059890747, + 1631.6800117492676, + 1595.6159830093384, + 1631.392002105713, + 1597.4080562591553, + 1634.2719793319702, + 1598.080039024353, + 1628.4799575805664, + 1596.8639850616455, + 1628.9279460906982, + 1592.6400423049927, + 1624.9279975891113, + 1598.3999967575073, + 1626.528024673462, + 1597.4080562591553, + 1626.431941986084, + 1595.6480503082275, + 1629.5039653778076, + 1597.7599620819092, + 1626.5599727630615, + 1595.744013786316, + 1629.5679807662964, + 1597.599983215332, + 1627.295970916748, + 1592.5120115280151, + 1627.8079748153687, + 1597.2479581832886, + 1626.5599727630615, + 1594.912052154541, + 1626.5599727630615, + 1597.6639986038208, + 1624.6720552444458, + 1593.34397315979, + 1631.168007850647, + 1599.6160507202148, + 1634.7520351409912, + 1595.8399772644043, + 1632.1280002593994, + 1596.1920022964478, + 1629.7279596328735, + 1596.127986907959, + 1634.6880197525024, + 1597.2479581832886, + 1625.8560419082642, + 1595.3279733657837, + 1628.607988357544, + 1591.1359786987305, + 1624.127984046936, + 1596.0960388183594, + 1631.8399906158447, + 1594.4960117340088, + 1629.7919750213623, + 1597.4400043487549, + 1633.1520080566406, + 1588.3519649505615, + 1627.8079748153687, + 1594.9759483337402, + 1628.8000345230103, + 1597.3440408706665, + 1631.1999559402466, + 1598.9439487457275, + 1630.079984664917, + 1594.3679809570312, + 1630.0159692764282, + 1597.5359678268433, + 1627.8400421142578 + ], + "candidate_backward": [ + 2241.5359020233154, + 2235.424041748047, + 2251.840114593506, + 2241.568088531494, + 2227.2000312805176, + 2238.464117050171, + 2228.224039077759, + 2242.464065551758, + 2256.927967071533, + 2230.272054672241, + 2261.023998260498, + 2233.344078063965, + 2233.2799434661865, + 2235.3599071502686, + 2231.328010559082, + 2228.2559871673584, + 2230.304002761841, + 2238.4960651397705, + 2238.368034362793, + 2229.2160987854004, + 2233.311891555786, + 2244.6720600128174, + 2234.4000339508057, + 2229.248046875, + 2241.663932800293, + 2233.4399223327637, + 2231.2960624694824, + 2236.4161014556885, + 2228.192090988159, + 2226.1440753936768, + 2232.3520183563232, + 2262.0480060577393, + 2238.431930541992, + 4604.896068572998, + 2221.0559844970703, + 2246.6559410095215, + 2212.8961086273193, + 2215.9359455108643, + 2212.671995162964, + 2228.2559871673584, + 2226.1760234832764, + 2227.231979370117, + 2210.848093032837, + 2221.0240364074707, + 2221.0559844970703, + 2210.911989212036, + 2453.5040855407715, + 2691.1680698394775, + 2235.424041748047, + 2221.08793258667, + 2219.007968902588, + 2220.223903656006, + 2210.848093032837, + 2210.848093032837, + 2217.9839611053467, + 2210.752010345459, + 2203.6800384521484, + 2214.911937713623, + 2210.7839584350586, + 2215.8079147338867, + 2210.815906524658, + 2211.872100830078, + 2201.6639709472656, + 2375.839948654175, + 2220.031976699829, + 2222.111940383911, + 2224.128007888794, + 2210.7839584350586, + 2210.815906524658, + 2221.08793258667, + 2210.7839584350586, + 2222.208023071289, + 2213.9201164245605, + 2211.8399143218994, + 2233.344078063965, + 2213.887929916382, + 2213.9201164245605, + 2223.1359481811523, + 2219.007968902588, + 2216.9599533081055, + 2215.8079147338867, + 2208.767890930176, + 2207.7760696411133, + 2212.8639221191406, + 2215.967893600464, + 2210.848093032837, + 2200.6399631500244, + 2313.0879402160645, + 2213.8240337371826, + 2210.848093032837, + 2214.9438858032227, + 2211.872100830078, + 2204.67209815979, + 2228.224039077759, + 2236.448049545288, + 2395.103931427002, + 2204.67209815979, + 2227.2000312805176, + 2222.048044204712, + 2210.848093032837, + 2209.7599506378174, + 2226.1440753936768, + 2212.8639221191406, + 2203.5839557647705, + 2201.472043991089, + 2210.7200622558594, + 2223.1040000915527, + 2214.911937713623, + 2212.8639221191406, + 2211.8399143218994, + 2207.8399658203125, + 2212.8961086273193, + 2207.7760696411133, + 2219.007968902588, + 2221.0240364074707, + 2222.0799922943115, + 2204.67209815979, + 2216.9599533081055, + 2205.6639194488525, + 2203.648090362549, + 2216.928005218506, + 2216.9599533081055, + 2212.8639221191406, + 2221.08793258667, + 2216.991901397705, + 2212.8639221191406, + 2222.0799922943115, + 2378.848075866699, + 2210.752010345459, + 2211.8079662323, + 2221.1201190948486, + 2223.1040000915527, + 2224.1599559783936, + 2209.728002548218, + 2213.1519317626953, + 2212.9600048065186, + 2211.9040489196777, + 2205.728054046631, + 2248.7359046936035, + 2219.935894012451, + 2205.6000232696533, + 2210.7839584350586, + 2211.8399143218994, + 2212.8639221191406, + 2203.648090362549, + 2234.4000339508057, + 2205.6961059570312, + 2206.7840099334717, + 2214.047908782959, + 2209.8240852355957, + 2209.8240852355957, + 2219.007968902588, + 2220.0000286102295, + 2206.752061843872, + 2212.8961086273193, + 2212.8639221191406, + 2208.928108215332, + 2215.9359455108643, + 2222.0799922943115, + 2310.175895690918, + 2216.032028198242, + 2208.672046661377, + 2218.0159091949463, + 2218.9760208129883, + 2210.7839584350586, + 2207.711935043335, + 2215.9359455108643, + 2210.815906524658, + 2206.7201137542725, + 2215.9359455108643, + 2229.2160987854004, + 2214.911937713623, + 2212.928056716919, + 2236.4161014556885, + 2228.2559871673584, + 2228.224039077759, + 2209.791898727417, + 2209.791898727417, + 2210.7839584350586, + 2206.7201137542725, + 2209.791898727417, + 2236.3839149475098, + 2217.9839611053467, + 2205.6961059570312, + 2209.791898727417, + 2207.8399658203125, + 2277.5039672851562, + 2301.9518852233887, + 2237.407922744751, + 2206.752061843872, + 2207.7438831329346, + 2213.7598991394043, + 2208.672046661377, + 2211.8079662323, + 2221.0559844970703, + 2209.8240852355957, + 2219.007968902588, + 2224.0960597991943, + 2211.8399143218994, + 2209.8560333251953 + ], + "provider_backward": [ + 8927.96802520752, + 8916.671752929688, + 8914.591789245605, + 8916.768074035645, + 8907.77587890625, + 8932.319641113281, + 8924.192428588867, + 8921.088218688965, + 8920.096397399902, + 8916.9921875, + 8923.871994018555, + 8925.919532775879, + 8920.096397399902, + 8914.624214172363, + 8909.567832946777, + 8917.759895324707, + 8909.88826751709, + 8911.871910095215, + 8917.023658752441, + 8908.831596374512, + 8913.920402526855, + 8914.87979888916, + 8914.624214172363, + 8919.584274291992, + 8918.815612792969, + 8912.704467773438, + 8919.103622436523, + 8930.047988891602, + 8921.088218688965, + 8902.655601501465, + 8902.655601501465, + 8918.08032989502, + 8916.000366210938, + 8908.896446228027, + 8923.839569091797, + 8923.935890197754, + 8918.68782043457, + 8914.591789245605, + 8903.29647064209, + 8911.552429199219, + 8919.103622436523, + 8919.136047363281, + 8921.119689941406, + 8909.82437133789, + 8905.535697937012, + 8908.86402130127, + 8914.591789245605, + 8905.664443969727, + 8910.55965423584, + 8920.928001403809, + 8909.855842590332, + 8910.49575805664, + 8907.615661621094, + 8919.072151184082, + 8902.688026428223, + 8906.75163269043, + 8916.9921875, + 8916.9921875, + 8918.848037719727, + 8914.655685424805, + 8929.311752319336, + 8918.047904968262, + 8912.60814666748, + 8927.807807922363, + 8924.192428588867, + 8910.880088806152, + 8916.000366210938, + 8906.784057617188, + 8925.18424987793, + 8922.975540161133, + 8910.4642868042, + 8905.792236328125, + 8915.040016174316, + 8912.896156311035, + 8932.255744934082, + 8907.487869262695, + 8911.840438842773, + 8912.863731384277, + 8926.015853881836, + 8907.77587890625, + 8909.855842590332, + 8915.103912353516, + 8915.040016174316, + 8905.599594116211, + 8925.919532775879, + 8926.239967346191, + 8916.768074035645, + 8901.472091674805, + 8916.9921875, + 8896.512031555176, + 8928.288459777832, + 8910.911560058594, + 8915.96794128418, + 8923.295974731445, + 8914.591789245605, + 8923.744201660156, + 8917.66357421875, + 8920.831680297852, + 8913.663864135742, + 8912.544250488281, + 8923.135757446289, + 8912.927627563477, + 8917.023658752441, + 8912.863731384277, + 8916.9921875, + 8921.343803405762, + 8911.552429199219, + 8916.576385498047, + 8923.808097839355, + 8917.792320251465, + 8918.08032989502, + 8914.688110351562, + 8934.399604797363, + 8921.088218688965, + 8919.039726257324, + 8902.688026428223, + 8926.239967346191, + 8920.096397399902, + 8907.872200012207, + 8919.072151184082, + 8912.960052490234, + 8918.784141540527, + 8921.82445526123, + 8916.640281677246, + 8927.359580993652, + 8907.77587890625, + 8915.807723999023, + 8919.008255004883, + 8911.904335021973, + 8921.152114868164, + 8921.152114868164, + 8912.863731384277, + 8926.176071166992, + 8906.335830688477, + 8914.976119995117, + 8921.055793762207, + 8907.808303833008, + 8916.000366210938, + 8916.000366210938, + 8914.976119995117, + 8923.135757446289, + 8910.880088806152, + 8912.927627563477, + 8911.775588989258, + 8923.839569091797, + 8914.624214172363, + 8919.136047363281, + 8910.880088806152, + 8927.103996276855, + 8930.335998535156, + 8913.984298706055, + 8910.880088806152, + 8907.71198272705, + 8915.007591247559, + 8916.031837463379, + 8907.83977508545, + 8918.880462646484, + 8919.072151184082, + 8927.295684814453, + 8911.520004272461, + 8912.67204284668, + 8919.039726257324, + 8912.927627563477, + 8931.360244750977, + 8907.83977508545, + 8922.143936157227, + 8919.072151184082, + 8941.472053527832, + 8931.136131286621, + 8914.527893066406, + 8923.935890197754, + 8914.94369506836, + 8919.872283935547, + 8912.927627563477, + 8912.896156311035, + 8906.784057617188, + 8918.047904968262, + 8914.015769958496, + 8903.679847717285, + 8911.616325378418, + 8908.608436584473, + 8912.639617919922, + 8921.82445526123, + 8916.671752929688, + 8910.752296447754, + 8900.383949279785, + 8917.0560836792, + 8922.176361083984, + 8932.38353729248, + 8912.768363952637, + 8918.01643371582, + 8921.119689941406, + 8920.80020904541, + 8914.78443145752, + 8916.671752929688, + 8916.447639465332, + 8906.784057617188, + 8904.767990112305, + 8933.4077835083, + 8923.135757446289 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 881.5039992332458, + "candidate_gbps": 4796.187465601401, + "candidate_peak_mib": 1344.0, + "provider_us": 1613.7920022010803, + "provider_gbps": 2619.8285939164075, + "provider_peak_mib": 2688.0, + "candidate_backward_us": 2216.480016708374, + "candidate_backward_gbps": 1907.46517005764, + "candidate_backward_peak_mib": 2688.64599609375, + "provider_backward_us": 8916.016101837158, + "provider_backward_gbps": 474.1869444503184, + "provider_backward_peak_mib": 5376.0 + } + ] +} diff --git a/reports/experiments/h3-adaln-projection-b200/figure.png b/reports/experiments/h3-adaln-projection-b200/figure.png new file mode 100644 index 000000000..a367fbf24 Binary files /dev/null and b/reports/experiments/h3-adaln-projection-b200/figure.png differ diff --git a/reports/experiments/h3-adaln-projection-b200/report.json b/reports/experiments/h3-adaln-projection-b200/report.json new file mode 100644 index 000000000..7f7bd36ad --- /dev/null +++ b/reports/experiments/h3-adaln-projection-b200/report.json @@ -0,0 +1,141 @@ +{ + "kind": "h3_operator_report", + "op": "adaln_projection_3mod", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "d06e1203b1001bb5055ae84a6ca816346227a94c", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "draws": 20, + "num_timesteps": 4, + "cuda_correctly_rounded": [ + 0.9998268485069275, + 0.9997907280921936, + 0.9998682141304016, + 0.9997829794883728, + 0.9998424053192139, + 0.9998501539230347, + 0.999793291091919, + 0.9998165369033813, + 0.9998036026954651, + 0.9998242855072021, + 0.9998217225074768, + 0.9997907280921936, + 0.9998268485069275, + 0.9998190999031067, + 0.9997622966766357, + 0.9998087882995605, + 0.9998397827148438, + 0.999832034111023, + 0.9998036026954651, + 0.9997958540916443 + ], + "provider_correctly_rounded": [ + 0.9982122182846069, + 0.9983181357383728, + 0.998591959476471, + 0.9982870817184448, + 0.9984291791915894, + 0.998426616191864, + 0.9982483386993408, + 0.9984395503997803, + 0.9984395503997803, + 0.9983981847763062, + 0.9984188675880432, + 0.9981166124343872, + 0.9981191754341125, + 0.9985195994377136, + 0.9978711605072021, + 0.9984421133995056, + 0.9985661506652832, + 0.9984757304191589, + 0.9983077645301819, + 0.9985015392303467 + ], + "early_cast_golden_match": [ + 0.45887067914009094, + 0.4315708577632904, + 0.4772522747516632, + 0.4299794137477875, + 0.4281528890132904, + 0.44502830505371094, + 0.4139591455459595, + 0.38746020197868347, + 0.40484198927879333, + 0.42634186148643494, + 0.40324538946151733, + 0.411001056432724, + 0.38042792677879333, + 0.4611622393131256, + 0.44856253266334534, + 0.46341246366500854, + 0.47497880458831787, + 0.4080067574977875, + 0.39637067914009094, + 0.39352625608444214 + ], + "rows_batch_invariant": true + }, + "perf": [ + { + "op": "adaln_projection_3mod", + "case": "T=1", + "backend": "H3AdaLNProjectionCudaOp", + "bytes": 520224768, + "candidate_us": 104.76800054311752, + "candidate_gbps": 4965.492949213059, + "candidate_peak_mib": 0.18994140625, + "provider_us": 114.73600193858147, + "provider_gbps": 4534.102280106274, + "provider_peak_mib": 0.18994140625 + }, + { + "op": "adaln_projection_3mod", + "case": "T=2", + "backend": "H3AdaLNProjectionCudaOp", + "bytes": 520224768, + "candidate_us": 104.68799993395805, + "candidate_gbps": 4969.287485940905, + "candidate_peak_mib": 0.37939453125, + "provider_us": 102.73599997162819, + "provider_gbps": 5063.704720289543, + "provider_peak_mib": 0.37939453125 + }, + { + "op": "adaln_projection_3mod", + "case": "T=3", + "backend": "H3AdaLNProjectionCudaOp", + "bytes": 520224768, + "candidate_us": 104.96000200510025, + "candidate_gbps": 4956.409661412937, + "candidate_peak_mib": 0.5693359375, + "provider_us": 103.15200313925743, + "provider_gbps": 5043.283234138316, + "provider_peak_mib": 0.5693359375 + }, + { + "op": "adaln_projection_3mod", + "case": "T=4", + "backend": "H3AdaLNProjectionCudaOp", + "bytes": 520224768, + "candidate_us": 105.55199906229973, + "candidate_gbps": 4928.611230687813, + "candidate_peak_mib": 0.7587890625, + "provider_us": 103.61599922180176, + "provider_gbps": 5020.699234742698, + "provider_peak_mib": 0.7587890625 + } + ] +} diff --git a/reports/experiments/h3-adaln-row-gather-b200/chain_replay.json b/reports/experiments/h3-adaln-row-gather-b200/chain_replay.json new file mode 100644 index 000000000..23f14359c --- /dev/null +++ b/reports/experiments/h3-adaln-row-gather-b200/chain_replay.json @@ -0,0 +1,3955 @@ +{ + "kind": "h3_conditioning_chain_replay", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "weights_sha256": { + "diffusion_pytorch_model-00001-of-00014.safetensors": { + "sha256": "2d847200c45c09dd7f973c1b096663068408ef851ee0b3711d059b6dc5dcd028", + "size_bytes": 4825958704 + } + }, + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "fa551c2aa0ee038be2e486a0ef5ff0fb03de49c0", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "stages": [ + "timestep_sinusoid_h3", + "timestep_mlp_fp32", + "adaln_projection_3mod", + "adaln_row_gather" + ], + "cases": [ + { + "num_timesteps": 1, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003897547721862793, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003897547721862793, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + } + ], + "backward_cases": [ + { + "num_timesteps": 1, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 14.676741600036621, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.5142298936843872 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.5142298936843872 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.020853519439697266, + "max_abs_vs_golden_over_absmax": 0.0014208548236638308, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 14.676741600036621, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.02845982275903225 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.02845982275903225 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.020853519439697266, + "max_abs_vs_golden_over_absmax": 0.0014208548236638308, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 20.589889526367188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.147125244140625e-05, + "max_abs_vs_golden_over_absmax": 1.5284808796423022e-06, + "correctly_rounded_fraction": 0.016372647136449814 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.147125244140625e-05, + "max_abs_vs_golden_over_absmax": 1.5284808796423022e-06, + "correctly_rounded_fraction": 0.016372647136449814 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.066619873046875, + "max_abs_vs_golden_over_absmax": 0.0032355624716728926, + "correctly_rounded_fraction": 2.4081899027805775e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 73.953369140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000110626220703125, + "max_abs_vs_golden_over_absmax": 1.4958915244278614e-06, + "correctly_rounded_fraction": 0.0524553582072258 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000110626220703125, + "max_abs_vs_golden_over_absmax": 1.4958915244278614e-06, + "correctly_rounded_fraction": 0.0524553582072258 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.23928070068359375, + "max_abs_vs_golden_over_absmax": 0.0032355617731809616, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 1.140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 5.351027357392013e-05, + "correctly_rounded_fraction": 0.9957404732704163 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 5.351027357392013e-05, + "correctly_rounded_fraction": 0.9957404732704163 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 2.6755136786960065e-05, + "correctly_rounded_fraction": 0.998334527015686 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.5096726417541504 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.3018841743469238, + "max_abs_vs_golden_over_absmax": 0.011344562284648418, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.0193452388048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.3018841743469238, + "max_abs_vs_golden_over_absmax": 0.011344562284648418, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 174.07337951660156, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.22721099853515625, + "max_abs_vs_golden_over_absmax": 0.0013052598806098104, + "correctly_rounded_fraction": 1.1487342817417812e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00023651123046875, + "max_abs_vs_golden_over_absmax": 1.358686972707801e-06, + "correctly_rounded_fraction": 0.013279437087476254 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.529754638671875, + "max_abs_vs_golden_over_absmax": 0.008787986822426319, + "correctly_rounded_fraction": 3.5292437132738996e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 625.2249755859375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8160858154296875, + "max_abs_vs_golden_over_absmax": 0.001305267447605729, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.3666955283042626e-06, + "correctly_rounded_fraction": 0.0372023805975914 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.4945068359375, + "max_abs_vs_golden_over_absmax": 0.008788047358393669, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 11.4375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.005464480724185705, + "correctly_rounded_fraction": 0.7395860552787781 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 4.269125565770082e-05, + "correctly_rounded_fraction": 0.9957284331321716 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.02185792289674282, + "correctly_rounded_fraction": 0.17165133357048035 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.875, + "max_abs_vs_golden_over_absmax": 0.019662922248244286, + "correctly_rounded_fraction": 0.17448949813842773 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.5178571343421936 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 10.64801025390625, + "max_abs_vs_golden_over_absmax": 0.019956601783633232, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.0357142873108387 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 10.64801025390625, + "max_abs_vs_golden_over_absmax": 0.019956601783633232, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 627.8563232421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.204010009765625, + "max_abs_vs_golden_over_absmax": 0.0019176520872861147, + "correctly_rounded_fraction": 2.560431767051341e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.360967758046172e-06, + "correctly_rounded_fraction": 0.015629082918167114 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 24.599517822265625, + "max_abs_vs_golden_over_absmax": 0.039180170744657516, + "correctly_rounded_fraction": 1.1764145710912999e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2255.091796875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.324493408203125, + "max_abs_vs_golden_over_absmax": 0.001917657325975597, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.3532743423638749e-06, + "correctly_rounded_fraction": 0.0520833358168602 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 88.3548583984375, + "max_abs_vs_golden_over_absmax": 0.03918015956878662, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 39.25, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.006369426846504211, + "correctly_rounded_fraction": 0.7380850911140442 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 4.976114723831415e-05, + "correctly_rounded_fraction": 0.9957199096679688 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.875, + "max_abs_vs_golden_over_absmax": 0.09872611612081528, + "correctly_rounded_fraction": 0.045190975069999695 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 15.0, + "max_abs_vs_golden_over_absmax": 0.09803921729326248, + "correctly_rounded_fraction": 0.04515955597162247 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 1096.3087158203125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.691497802734375, + "max_abs_vs_golden_over_absmax": 0.00154290278442204, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0022430419921875, + "max_abs_vs_golden_over_absmax": 2.045994961008546e-06, + "correctly_rounded_fraction": 0.5103236436843872 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 129.3043670654297, + "max_abs_vs_golden_over_absmax": 0.11794521659612656, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 1096.3087158203125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.691497802734375, + "max_abs_vs_golden_over_absmax": 0.00154290278442204, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0022430419921875, + "max_abs_vs_golden_over_absmax": 2.045994961008546e-06, + "correctly_rounded_fraction": 0.0206473208963871 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 129.3043670654297, + "max_abs_vs_golden_over_absmax": 0.11794521659612656, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2459.76708984375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.49169921875, + "max_abs_vs_golden_over_absmax": 0.0010129817528650165, + "correctly_rounded_fraction": 5.453027551993728e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0032958984375, + "max_abs_vs_golden_over_absmax": 1.3399229601418483e-06, + "correctly_rounded_fraction": 0.016593534499406815 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 322.1355895996094, + "max_abs_vs_golden_over_absmax": 0.1309618204832077, + "correctly_rounded_fraction": 6.920085837691659e-08 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8834.82421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.94921875, + "max_abs_vs_golden_over_absmax": 0.0010129481088370085, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.01171875, + "max_abs_vs_golden_over_absmax": 1.3264270819490775e-06, + "correctly_rounded_fraction": 0.063988097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1157.0247802734375, + "max_abs_vs_golden_over_absmax": 0.1309618353843689, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 121.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004115226212888956, + "correctly_rounded_fraction": 0.737601637840271 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.0005144032766111195, + "correctly_rounded_fraction": 0.9956783652305603 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 43.5, + "max_abs_vs_golden_over_absmax": 0.3580246865749359, + "correctly_rounded_fraction": 0.01575867086648941 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0078125, + "max_abs_vs_golden_over_absmax": 1.641281596675981e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0078125, + "max_abs_vs_golden_over_absmax": 1.641281596675981e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 170.0, + "max_abs_vs_golden_over_absmax": 0.3571428656578064, + "correctly_rounded_fraction": 0.015666335821151733 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024131298065185547, + "max_abs_vs_golden_over_absmax": 0.0027526128105819225, + "correctly_rounded_fraction": 2.1798271063744323e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024131298065185547, + "max_abs_vs_golden_over_absmax": 0.0027526128105819225, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 21.970394134521484, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04718971252441406, + "max_abs_vs_golden_over_absmax": 0.0021478773560374975, + "correctly_rounded_fraction": 1.972224526980426e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 78.91429901123047, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.16949462890625, + "max_abs_vs_golden_over_absmax": 0.0021478317212313414, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 3.578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 8.528930266038515e-06, + "correctly_rounded_fraction": 0.998363733291626 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 95.00251007080078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.126617431640625, + "max_abs_vs_golden_over_absmax": 0.0013327798806130886, + "correctly_rounded_fraction": 5.158923886483535e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00016021728515625, + "max_abs_vs_golden_over_absmax": 1.6864531744431588e-06, + "correctly_rounded_fraction": 0.02457391656935215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.3639969825744629, + "max_abs_vs_golden_over_absmax": 0.003831445937976241, + "correctly_rounded_fraction": 7.266090165103378e-07 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 95.00251007080078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.12661361694335938, + "max_abs_vs_golden_over_absmax": 0.0013327397173270583, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00012969970703125, + "max_abs_vs_golden_over_absmax": 1.365223965876794e-06, + "correctly_rounded_fraction": 0.02752976305782795 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.36399686336517334, + "max_abs_vs_golden_over_absmax": 0.003831444773823023, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.80873107910156, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2622566223144531, + "max_abs_vs_golden_over_absmax": 0.001508880639448762, + "correctly_rounded_fraction": 2.3874295948189683e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00018310546875, + "max_abs_vs_golden_over_absmax": 1.0534882903812104e-06, + "correctly_rounded_fraction": 0.014370596036314964 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.4940261840820312, + "max_abs_vs_golden_over_absmax": 0.008595806546509266, + "correctly_rounded_fraction": 2.4912308163038688e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.3004150390625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.941986083984375, + "max_abs_vs_golden_over_absmax": 0.0015088666696101427, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00067138671875, + "max_abs_vs_golden_over_absmax": 1.0754224604170304e-06, + "correctly_rounded_fraction": 0.0457589291036129 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.366363525390625, + "max_abs_vs_golden_over_absmax": 0.008595802821218967, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 25.75, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.004854368977248669, + "correctly_rounded_fraction": 0.6666338443756104 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.8962378817377612e-05, + "correctly_rounded_fraction": 0.9989553093910217 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.019417475908994675, + "correctly_rounded_fraction": 0.2250475138425827 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.6273354887962341 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.625, + "max_abs_vs_golden_over_absmax": 0.014044944196939468, + "correctly_rounded_fraction": 0.21870866417884827 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 527.6658935546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.27825927734375, + "max_abs_vs_golden_over_absmax": 0.0005273398710414767, + "correctly_rounded_fraction": 4.359654212748865e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00042724609375, + "max_abs_vs_golden_over_absmax": 8.096905617094308e-07, + "correctly_rounded_fraction": 0.0324409119784832 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 12.919952392578125, + "max_abs_vs_golden_over_absmax": 0.02448510006070137, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 527.6620483398438, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2781982421875, + "max_abs_vs_golden_over_absmax": 0.000527228054124862, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00042724609375, + "max_abs_vs_golden_over_absmax": 8.096964734249923e-07, + "correctly_rounded_fraction": 0.0355282761156559 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 12.919921875, + "max_abs_vs_golden_over_absmax": 0.02448522113263607, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 630.265869140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8614501953125, + "max_abs_vs_golden_over_absmax": 0.001366804470308125, + "correctly_rounded_fraction": 2.255947947560344e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000732421875, + "max_abs_vs_golden_over_absmax": 1.162084004135977e-06, + "correctly_rounded_fraction": 0.01656709983944893 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 19.454803466796875, + "max_abs_vs_golden_over_absmax": 0.03086761385202408, + "correctly_rounded_fraction": 1.3148163588994066e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2263.81787109375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.09423828125, + "max_abs_vs_golden_over_absmax": 0.0013668229803442955, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0025634765625, + "max_abs_vs_golden_over_absmax": 1.1323687658659765e-06, + "correctly_rounded_fraction": 0.0699404776096344 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 69.87863159179688, + "max_abs_vs_golden_over_absmax": 0.030867602676153183, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 115.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004347825888544321, + "correctly_rounded_fraction": 0.6659318208694458 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.03125, + "max_abs_vs_golden_over_absmax": 0.00027173911803402007, + "correctly_rounded_fraction": 0.9989242553710938 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.5, + "max_abs_vs_golden_over_absmax": 0.04782608523964882, + "correctly_rounded_fraction": 0.06274387985467911 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6259506940841675 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 8.25, + "max_abs_vs_golden_over_absmax": 0.053921569138765335, + "correctly_rounded_fraction": 0.0617869533598423 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 744.156982421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.7833938598632812, + "max_abs_vs_golden_over_absmax": 0.0023965290747582912, + "correctly_rounded_fraction": 1.2352353223832324e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001129150390625, + "max_abs_vs_golden_over_absmax": 1.5173551446423517e-06, + "correctly_rounded_fraction": 0.0301063172519207 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 87.67558288574219, + "max_abs_vs_golden_over_absmax": 0.1178186684846878, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 744.156982421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.7833938598632812, + "max_abs_vs_golden_over_absmax": 0.0023965290747582912, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00109100341796875, + "max_abs_vs_golden_over_absmax": 1.4660930673926487e-06, + "correctly_rounded_fraction": 0.0329241082072258 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 87.67555236816406, + "max_abs_vs_golden_over_absmax": 0.11781862378120422, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2462.752685546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.5404052734375, + "max_abs_vs_golden_over_absmax": 0.001437580562196672, + "correctly_rounded_fraction": 2.7334339392837137e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.002685546875, + "max_abs_vs_golden_over_absmax": 1.0904655027843546e-06, + "correctly_rounded_fraction": 0.0176222063601017 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 200.62046813964844, + "max_abs_vs_golden_over_absmax": 0.08146188408136368, + "correctly_rounded_fraction": 6.920085979800206e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8845.9306640625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 12.716552734375, + "max_abs_vs_golden_over_absmax": 0.0014375596074387431, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.009765625, + "max_abs_vs_golden_over_absmax": 1.103968088500551e-06, + "correctly_rounded_fraction": 0.0658482164144516 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 720.5938720703125, + "max_abs_vs_golden_over_absmax": 0.08146049082279205, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 298.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.003355704713612795, + "correctly_rounded_fraction": 0.6657332181930542 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0016778523568063974, + "correctly_rounded_fraction": 0.9988868236541748 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 47.0, + "max_abs_vs_golden_over_absmax": 0.15771812200546265, + "correctly_rounded_fraction": 0.022447003051638603 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.6236565709114075 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0010504202218726277, + "correctly_rounded_fraction": 0.9999483227729797 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 62.25, + "max_abs_vs_golden_over_absmax": 0.13077731430530548, + "correctly_rounded_fraction": 0.022641781717538834 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 11.570573806762695, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.2099742889404297e-05, + "max_abs_vs_golden_over_absmax": 1.0457340522407321e-06, + "correctly_rounded_fraction": 0.0445265993475914 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.2099742889404297e-05, + "max_abs_vs_golden_over_absmax": 1.0457340522407321e-06, + "correctly_rounded_fraction": 0.0445265993475914 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.025829792022705078, + "max_abs_vs_golden_over_absmax": 0.0022323690354824066, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 11.57003402709961, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.208484172821045e-05, + "max_abs_vs_golden_over_absmax": 1.0444948657095665e-06, + "correctly_rounded_fraction": 0.0403645858168602 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.208484172821045e-05, + "max_abs_vs_golden_over_absmax": 1.0444948657095665e-06, + "correctly_rounded_fraction": 0.0403645858168602 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.025829792022705078, + "max_abs_vs_golden_over_absmax": 0.0022324733436107635, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 20.580326080322266, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.52587890625e-05, + "max_abs_vs_golden_over_absmax": 7.414259926008526e-07, + "correctly_rounded_fraction": 0.019046567380428314 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.52587890625e-05, + "max_abs_vs_golden_over_absmax": 7.414259926008526e-07, + "correctly_rounded_fraction": 0.019046567380428314 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04852867126464844, + "max_abs_vs_golden_over_absmax": 0.002358012832701206, + "correctly_rounded_fraction": 2.159066752938088e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 73.91921997070312, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.7220458984375e-05, + "max_abs_vs_golden_over_absmax": 7.740944738543476e-07, + "correctly_rounded_fraction": 0.1011904776096344 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.7220458984375e-05, + "max_abs_vs_golden_over_absmax": 7.740944738543476e-07, + "correctly_rounded_fraction": 0.1011904776096344 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1743011474609375, + "max_abs_vs_golden_over_absmax": 0.002357994904741645, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 2.859375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 2.134562782885041e-05, + "correctly_rounded_fraction": 0.9968344569206238 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 2.134562782885041e-05, + "correctly_rounded_fraction": 0.9968344569206238 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 1.0672813914425205e-05, + "correctly_rounded_fraction": 0.9970113635063171 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087639331817627, + "max_abs_vs_golden_over_absmax": 0.0014897150686010718, + "correctly_rounded_fraction": 4.359654212748865e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00018215179443359375, + "max_abs_vs_golden_over_absmax": 2.494892214599531e-06, + "correctly_rounded_fraction": 0.02586873434484005 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.2452354431152344, + "max_abs_vs_golden_over_absmax": 0.017055708914995193, + "correctly_rounded_fraction": 1.52587890625e-05 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087636947631836, + "max_abs_vs_golden_over_absmax": 0.0014897118089720607, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0001819133758544922, + "max_abs_vs_golden_over_absmax": 2.4916266738728154e-06, + "correctly_rounded_fraction": 0.0232514888048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.2452354431152344, + "max_abs_vs_golden_over_absmax": 0.017055708914995193, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.8400421142578, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2456045150756836, + "max_abs_vs_golden_over_absmax": 0.001412819023244083, + "correctly_rounded_fraction": 2.089865847665351e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000171661376953125, + "max_abs_vs_golden_over_absmax": 9.874673878584872e-07, + "correctly_rounded_fraction": 0.012266821227967739 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.2523880004882812, + "max_abs_vs_golden_over_absmax": 0.00720425508916378, + "correctly_rounded_fraction": 3.3216410884051584e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.4014892578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8821792602539062, + "max_abs_vs_golden_over_absmax": 0.0014128397451713681, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0006103515625, + "max_abs_vs_golden_over_absmax": 9.774985301191919e-07, + "correctly_rounded_fraction": 0.0412946417927742 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 4.498443603515625, + "max_abs_vs_golden_over_absmax": 0.007204408757388592, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 27.375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.004566209856420755, + "correctly_rounded_fraction": 0.6352493762969971 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.7836757251643576e-05, + "correctly_rounded_fraction": 0.9988961815834045 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.3125, + "max_abs_vs_golden_over_absmax": 0.01141552533954382, + "correctly_rounded_fraction": 0.26146775484085083 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.6106977462768555 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.4375, + "max_abs_vs_golden_over_absmax": 0.009831461124122143, + "correctly_rounded_fraction": 0.25494998693466187 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 216.79254150390625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.34185791015625, + "max_abs_vs_golden_over_absmax": 0.0015768896555528045, + "correctly_rounded_fraction": 4.650297705666162e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.293441634217743e-06, + "correctly_rounded_fraction": 0.02635483630001545 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 7.714607238769531, + "max_abs_vs_golden_over_absmax": 0.035585206001996994, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 189.4628448486328, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.3418426513671875, + "max_abs_vs_golden_over_absmax": 0.0018042727606371045, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.912763415632071e-06, + "correctly_rounded_fraction": 0.0204613097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 6.5852203369140625, + "max_abs_vs_golden_over_absmax": 0.03475731611251831, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 632.3684692382812, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.124969482421875, + "max_abs_vs_golden_over_absmax": 0.0017789778066799045, + "correctly_rounded_fraction": 2.3805096134310588e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009765625, + "max_abs_vs_golden_over_absmax": 1.544293468214164e-06, + "correctly_rounded_fraction": 0.011527756229043007 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 15.004081726074219, + "max_abs_vs_golden_over_absmax": 0.023726802319288254, + "correctly_rounded_fraction": 1.7300214949500514e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2271.419677734375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.0406494140625, + "max_abs_vs_golden_over_absmax": 0.0017789092380553484, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0037841796875, + "max_abs_vs_golden_over_absmax": 1.6659976154187461e-06, + "correctly_rounded_fraction": 0.0338541679084301 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 53.892181396484375, + "max_abs_vs_golden_over_absmax": 0.0237262099981308, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 94.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.005291005130857229, + "correctly_rounded_fraction": 0.6358715295791626 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 4.13359775848221e-05, + "correctly_rounded_fraction": 0.9988698363304138 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.25, + "max_abs_vs_golden_over_absmax": 0.0555555559694767, + "correctly_rounded_fraction": 0.07476150244474411 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6086929440498352 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 2.5531046048854478e-05, + "correctly_rounded_fraction": 0.9999896287918091 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 7.0, + "max_abs_vs_golden_over_absmax": 0.04575163498520851, + "correctly_rounded_fraction": 0.074373759329319 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 1221.45751953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.3370361328125, + "max_abs_vs_golden_over_absmax": 0.0010946234688162804, + "correctly_rounded_fraction": 6.5394810917496216e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001739501953125, + "max_abs_vs_golden_over_absmax": 1.424119886905828e-06, + "correctly_rounded_fraction": 0.0268104188144207 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 42.02140808105469, + "max_abs_vs_golden_over_absmax": 0.03440267592668533, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 1221.45751953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.3369140625, + "max_abs_vs_golden_over_absmax": 0.0010945235844701529, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001708984375, + "max_abs_vs_golden_over_absmax": 1.3991353853270994e-06, + "correctly_rounded_fraction": 0.032738097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 42.016990661621094, + "max_abs_vs_golden_over_absmax": 0.03439905866980553, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2460.399169921875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.100341796875, + "max_abs_vs_golden_over_absmax": 0.0012600970221683383, + "correctly_rounded_fraction": 2.9894770705141127e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00274658203125, + "max_abs_vs_golden_over_absmax": 1.1163156159454957e-06, + "correctly_rounded_fraction": 0.01790323108434677 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 143.638427734375, + "max_abs_vs_golden_over_absmax": 0.0583801306784153, + "correctly_rounded_fraction": 4.152051360506448e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8837.2431640625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 11.13623046875, + "max_abs_vs_golden_over_absmax": 0.0012601475464180112, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00927734375, + "max_abs_vs_golden_over_absmax": 1.0498006304260343e-06, + "correctly_rounded_fraction": 0.0632440522313118 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 515.924560546875, + "max_abs_vs_golden_over_absmax": 0.05838071182370186, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 288.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.0069444444961845875, + "correctly_rounded_fraction": 0.6357954144477844 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0017361111240461469, + "correctly_rounded_fraction": 0.9988410472869873 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 34.25, + "max_abs_vs_golden_over_absmax": 0.1189236119389534, + "correctly_rounded_fraction": 0.02671528421342373 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.6103463768959045 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.00026260505546815693, + "correctly_rounded_fraction": 0.9999379515647888 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 52.0, + "max_abs_vs_golden_over_absmax": 0.10924369841814041, + "correctly_rounded_fraction": 0.025752313435077667 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024135589599609375, + "max_abs_vs_golden_over_absmax": 0.002753102220594883, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024135589599609375, + "max_abs_vs_golden_over_absmax": 0.002753102220594883, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 21.970394134521484, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04718971252441406, + "max_abs_vs_golden_over_absmax": 0.0021478773560374975, + "correctly_rounded_fraction": 1.972224526980426e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 78.91429901123047, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.16949462890625, + "max_abs_vs_golden_over_absmax": 0.0021478317212313414, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 3.578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 8.528930266038515e-06, + "correctly_rounded_fraction": 0.998363733291626 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 54.46367263793945, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.11591720581054688, + "max_abs_vs_golden_over_absmax": 0.0021283398382365704, + "correctly_rounded_fraction": 3.6330450257082703e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 1.1906992085641832e-06, + "correctly_rounded_fraction": 0.02683221735060215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5249824523925781, + "max_abs_vs_golden_over_absmax": 0.009639130905270576, + "correctly_rounded_fraction": 7.2660900514165405e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 54.46367263793945, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.11591720581054688, + "max_abs_vs_golden_over_absmax": 0.0021283398382365704, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.340576171875e-05, + "max_abs_vs_golden_over_absmax": 9.805758054426406e-07, + "correctly_rounded_fraction": 0.0271577388048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5249824523925781, + "max_abs_vs_golden_over_absmax": 0.009639130905270576, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.8180694580078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2628202438354492, + "max_abs_vs_golden_over_absmax": 0.001512042130343616, + "correctly_rounded_fraction": 1.4670581549580675e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00014495849609375, + "max_abs_vs_golden_over_absmax": 8.339667942891538e-07, + "correctly_rounded_fraction": 0.012976337224245071 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.9543752670288086, + "max_abs_vs_golden_over_absmax": 0.0054906560108065605, + "correctly_rounded_fraction": 5.2592654355976265e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.3176879882812, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.9439964294433594, + "max_abs_vs_golden_over_absmax": 0.0015120450407266617, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000537872314453125, + "max_abs_vs_golden_over_absmax": 8.615362503405777e-07, + "correctly_rounded_fraction": 0.0386904776096344 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.4279212951660156, + "max_abs_vs_golden_over_absmax": 0.0054906681180000305, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 25.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.0049019609577953815, + "correctly_rounded_fraction": 0.619565486907959 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.914828499138821e-05, + "correctly_rounded_fraction": 0.9988008141517639 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.375, + "max_abs_vs_golden_over_absmax": 0.014705882407724857, + "correctly_rounded_fraction": 0.2861616015434265 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.597811222076416 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.01123595517128706, + "correctly_rounded_fraction": 0.28086763620376587 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 326.7618713378906, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.541351318359375, + "max_abs_vs_golden_over_absmax": 0.0016567150596529245, + "correctly_rounded_fraction": 3.705705967149697e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000244140625, + "max_abs_vs_golden_over_absmax": 7.471514891221887e-07, + "correctly_rounded_fraction": 0.0328994020819664 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.4827728271484375, + "max_abs_vs_golden_over_absmax": 0.010658442974090576, + "correctly_rounded_fraction": 7.266090165103378e-07 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 326.7618713378906, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5413436889648438, + "max_abs_vs_golden_over_absmax": 0.0016566917765885592, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000213623046875, + "max_abs_vs_golden_over_absmax": 6.537575814036245e-07, + "correctly_rounded_fraction": 0.0319940485060215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.4827651977539062, + "max_abs_vs_golden_over_absmax": 0.01065841969102621, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 630.1971435546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.950373649597168, + "max_abs_vs_golden_over_absmax": 0.0015080576995387673, + "correctly_rounded_fraction": 2.4981509341159835e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0006103515625, + "max_abs_vs_golden_over_absmax": 9.685089708000305e-07, + "correctly_rounded_fraction": 0.013629386201500893 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.818405151367188, + "max_abs_vs_golden_over_absmax": 0.023513920605182648, + "correctly_rounded_fraction": 8.304102721012896e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2263.544677734375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.4134902954101562, + "max_abs_vs_golden_over_absmax": 0.0015080287121236324, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.002197265625, + "max_abs_vs_golden_over_absmax": 9.70718929238501e-07, + "correctly_rounded_fraction": 0.0498511902987957 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 53.2242431640625, + "max_abs_vs_golden_over_absmax": 0.02351367101073265, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 104.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004807692486792803, + "correctly_rounded_fraction": 0.6189784407615662 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.03125, + "max_abs_vs_golden_over_absmax": 0.0003004807804245502, + "correctly_rounded_fraction": 0.9987642765045166 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.0, + "max_abs_vs_golden_over_absmax": 0.028846153989434242, + "correctly_rounded_fraction": 0.0855673998594284 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6000640392303467 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 2.5531046048854478e-05, + "correctly_rounded_fraction": 0.9999793171882629 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.25, + "max_abs_vs_golden_over_absmax": 0.03431372717022896, + "correctly_rounded_fraction": 0.08447007089853287 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436309814453125, + "max_abs_vs_golden_over_absmax": 0.0007790509844198823, + "correctly_rounded_fraction": 3.4150623832829297e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.02686346136033535 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 38.62217712402344, + "max_abs_vs_golden_over_absmax": 0.046748287975788116, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436004638671875, + "max_abs_vs_golden_over_absmax": 0.0007790140807628632, + "correctly_rounded_fraction": 0.00018601190822664648 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.0282738097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 38.622161865234375, + "max_abs_vs_golden_over_absmax": 0.046748269349336624, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2460.831298828125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.702392578125, + "max_abs_vs_golden_over_absmax": 0.0010981624945998192, + "correctly_rounded_fraction": 3.480803206912242e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.2401328604028095e-06, + "correctly_rounded_fraction": 0.018226468935608864 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 101.59765625, + "max_abs_vs_golden_over_absmax": 0.041285909712314606, + "correctly_rounded_fraction": 8.99611166005343e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8838.8046875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 9.70703125, + "max_abs_vs_golden_over_absmax": 0.0010982289677485824, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.01171875, + "max_abs_vs_golden_over_absmax": 1.3258297713036882e-06, + "correctly_rounded_fraction": 0.0569196455180645 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 364.91552734375, + "max_abs_vs_golden_over_absmax": 0.041285619139671326, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 282.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.003546099178493023, + "correctly_rounded_fraction": 0.6200082898139954 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.00022163119865581393, + "correctly_rounded_fraction": 0.9987457394599915 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 25.0, + "max_abs_vs_golden_over_absmax": 0.08865248411893845, + "correctly_rounded_fraction": 0.030473006889224052 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.5982245802879333 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 8.206407983379904e-06, + "correctly_rounded_fraction": 0.9999586343765259 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 40.0, + "max_abs_vs_golden_over_absmax": 0.08403361588716507, + "correctly_rounded_fraction": 0.03078497014939785 + } + } + } + } + ] +} diff --git a/reports/experiments/h3-adaln-row-gather-b200/figure.png b/reports/experiments/h3-adaln-row-gather-b200/figure.png new file mode 100644 index 000000000..3cb869900 Binary files /dev/null and b/reports/experiments/h3-adaln-row-gather-b200/figure.png differ diff --git a/reports/experiments/h3-adaln-row-gather-b200/report.json b/reports/experiments/h3-adaln-row-gather-b200/report.json new file mode 100644 index 000000000..b8bbad0d7 --- /dev/null +++ b/reports/experiments/h3-adaln-row-gather-b200/report.json @@ -0,0 +1,804 @@ +{ + "kind": "h3_operator_report", + "op": "adaln_row_gather", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "80e46097302132c633dc6fc1a311097b616ac966", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "forward_bitwise_vs_index_select": { + "1": true, + "257": true, + "4097": true, + "32768": true + }, + "op_backward": { + "cuda": { + "repeat_bitwise_equal": true, + "correctly_rounded_fraction": 0.999996542930603 + }, + "provider": { + "repeat_bitwise_equal": false, + "correctly_rounded_fraction": 0.07806643843650818 + } + }, + "chain_backward": [ + { + "num_timesteps": 1, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.5096726417541504 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.7000911235809326, + "max_abs_vs_golden_over_absmax": 0.006100563798099756, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.0193452388048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.7000911235809326, + "max_abs_vs_golden_over_absmax": 0.006100563798099756, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 174.07337951660156, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.22721099853515625, + "max_abs_vs_golden_over_absmax": 0.0013052598806098104, + "correctly_rounded_fraction": 1.1487342817417812e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00023651123046875, + "max_abs_vs_golden_over_absmax": 1.358686972707801e-06, + "correctly_rounded_fraction": 0.013279437087476254 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.7896575927734375, + "max_abs_vs_golden_over_absmax": 0.010281052440404892, + "correctly_rounded_fraction": 2.076025680253224e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 625.2249755859375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8160858154296875, + "max_abs_vs_golden_over_absmax": 0.001305267447605729, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.3666955283042626e-06, + "correctly_rounded_fraction": 0.0372023805975914 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 6.427978515625, + "max_abs_vs_golden_over_absmax": 0.010281064547598362, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 11.4375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.005464480724185705, + "correctly_rounded_fraction": 0.7395860552787781 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 4.269125565770082e-05, + "correctly_rounded_fraction": 0.9957284331321716 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.02185792289674282, + "correctly_rounded_fraction": 0.17071296274662018 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.9375, + "max_abs_vs_golden_over_absmax": 0.021067416295409203, + "correctly_rounded_fraction": 0.17304272949695587 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087639331817627, + "max_abs_vs_golden_over_absmax": 0.0014897150686010718, + "correctly_rounded_fraction": 4.359654212748865e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00018215179443359375, + "max_abs_vs_golden_over_absmax": 2.494892214599531e-06, + "correctly_rounded_fraction": 0.02586873434484005 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.33492279052734375, + "max_abs_vs_golden_over_absmax": 0.004587362054735422, + "correctly_rounded_fraction": 2.9064360660413513e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087636947631836, + "max_abs_vs_golden_over_absmax": 0.0014897118089720607, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0001819133758544922, + "max_abs_vs_golden_over_absmax": 2.4916266738728154e-06, + "correctly_rounded_fraction": 0.0232514888048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.3349189758300781, + "max_abs_vs_golden_over_absmax": 0.004587309900671244, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.8400421142578, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2456045150756836, + "max_abs_vs_golden_over_absmax": 0.001412819023244083, + "correctly_rounded_fraction": 2.089865847665351e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000171661376953125, + "max_abs_vs_golden_over_absmax": 9.874673878584872e-07, + "correctly_rounded_fraction": 0.012266821227967739 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.0677642822265625, + "max_abs_vs_golden_over_absmax": 0.006142222788184881, + "correctly_rounded_fraction": 3.806047288890113e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.4014892578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8821792602539062, + "max_abs_vs_golden_over_absmax": 0.0014128397451713681, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0006103515625, + "max_abs_vs_golden_over_absmax": 9.774985301191919e-07, + "correctly_rounded_fraction": 0.0412946417927742 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.835296630859375, + "max_abs_vs_golden_over_absmax": 0.0061423564329743385, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 27.375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.004566209856420755, + "correctly_rounded_fraction": 0.6352493762969971 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.7836757251643576e-05, + "correctly_rounded_fraction": 0.9988961815834045 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00913241971284151, + "correctly_rounded_fraction": 0.2612312436103821 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.6106977462768555 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.01123595517128706, + "correctly_rounded_fraction": 0.25560101866722107 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.5178571343421936 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 15.295928955078125, + "max_abs_vs_golden_over_absmax": 0.02866777405142784, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.0357142873108387 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 15.295928955078125, + "max_abs_vs_golden_over_absmax": 0.02866777405142784, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 627.8563232421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.204010009765625, + "max_abs_vs_golden_over_absmax": 0.0019176520872861147, + "correctly_rounded_fraction": 2.560431767051341e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.360967758046172e-06, + "correctly_rounded_fraction": 0.015629082918167114 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 24.82110595703125, + "max_abs_vs_golden_over_absmax": 0.0395330972969532, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2255.091796875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.324493408203125, + "max_abs_vs_golden_over_absmax": 0.001917657325975597, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.3532743423638749e-06, + "correctly_rounded_fraction": 0.0520833358168602 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 89.1507568359375, + "max_abs_vs_golden_over_absmax": 0.0395330935716629, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 39.25, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.006369426846504211, + "correctly_rounded_fraction": 0.7380850911140442 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 4.976114723831415e-05, + "correctly_rounded_fraction": 0.9957199096679688 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.375, + "max_abs_vs_golden_over_absmax": 0.08598726242780685, + "correctly_rounded_fraction": 0.045540034770965576 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 13.0, + "max_abs_vs_golden_over_absmax": 0.08496732264757156, + "correctly_rounded_fraction": 0.0459449402987957 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 216.79254150390625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.34185791015625, + "max_abs_vs_golden_over_absmax": 0.0015768896555528045, + "correctly_rounded_fraction": 4.650297705666162e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.293441634217743e-06, + "correctly_rounded_fraction": 0.02635483630001545 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.418106079101562, + "max_abs_vs_golden_over_absmax": 0.06650646775960922, + "correctly_rounded_fraction": 2.9064360660413513e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 189.4628448486328, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.3418426513671875, + "max_abs_vs_golden_over_absmax": 0.0018042727606371045, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.912763415632071e-06, + "correctly_rounded_fraction": 0.0204613097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.4180908203125, + "max_abs_vs_golden_over_absmax": 0.07609983533620834, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 632.3684692382812, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.124969482421875, + "max_abs_vs_golden_over_absmax": 0.0017789778066799045, + "correctly_rounded_fraction": 2.3805096134310588e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009765625, + "max_abs_vs_golden_over_absmax": 1.544293468214164e-06, + "correctly_rounded_fraction": 0.011527756229043007 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.604507446289062, + "max_abs_vs_golden_over_absmax": 0.023094933480024338, + "correctly_rounded_fraction": 1.5224188700813102e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2271.419677734375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.0406494140625, + "max_abs_vs_golden_over_absmax": 0.0017789092380553484, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0037841796875, + "max_abs_vs_golden_over_absmax": 1.6659976154187461e-06, + "correctly_rounded_fraction": 0.0338541679084301 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 52.45654296875, + "max_abs_vs_golden_over_absmax": 0.023094166070222855, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 94.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.005291005130857229, + "correctly_rounded_fraction": 0.6358715295791626 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 4.13359775848221e-05, + "correctly_rounded_fraction": 0.9988698363304138 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 4.25, + "max_abs_vs_golden_over_absmax": 0.04497354477643967, + "correctly_rounded_fraction": 0.07502789795398712 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6086929440498352 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 2.5531046048854478e-05, + "correctly_rounded_fraction": 0.9999896287918091 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 7.0, + "max_abs_vs_golden_over_absmax": 0.04575163498520851, + "correctly_rounded_fraction": 0.07290633022785187 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436309814453125, + "max_abs_vs_golden_over_absmax": 0.0007790509844198823, + "correctly_rounded_fraction": 3.4150623832829297e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.02686346136033535 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 60.89818572998047, + "max_abs_vs_golden_over_absmax": 0.07371117174625397, + "correctly_rounded_fraction": 7.266090165103378e-07 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436004638671875, + "max_abs_vs_golden_over_absmax": 0.0007790140807628632, + "correctly_rounded_fraction": 0.00018601190822664648 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.0282738097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 60.89817810058594, + "max_abs_vs_golden_over_absmax": 0.07371116429567337, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2460.831298828125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.702392578125, + "max_abs_vs_golden_over_absmax": 0.0010981624945998192, + "correctly_rounded_fraction": 3.480803206912242e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.2401328604028095e-06, + "correctly_rounded_fraction": 0.018226468935608864 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 139.32177734375, + "max_abs_vs_golden_over_absmax": 0.056615736335515976, + "correctly_rounded_fraction": 1.3840171675383317e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8838.8046875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 9.70703125, + "max_abs_vs_golden_over_absmax": 0.0010982289677485824, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.01171875, + "max_abs_vs_golden_over_absmax": 1.3258297713036882e-06, + "correctly_rounded_fraction": 0.0569196455180645 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 500.4111328125, + "max_abs_vs_golden_over_absmax": 0.05661524832248688, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 282.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.003546099178493023, + "correctly_rounded_fraction": 0.6200082898139954 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.00022163119865581393, + "correctly_rounded_fraction": 0.9987457394599915 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 24.0, + "max_abs_vs_golden_over_absmax": 0.08510638028383255, + "correctly_rounded_fraction": 0.030594661831855774 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.5982245802879333 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 8.206407983379904e-06, + "correctly_rounded_fraction": 0.9999586343765259 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 38.0, + "max_abs_vs_golden_over_absmax": 0.07983193546533585, + "correctly_rounded_fraction": 0.029761902987957 + } + } + } + } + ] + }, + "perf": [ + { + "op": "adaln_row_gather", + "case": "S=4097", + "backend": "H3AdaLNRowGatherCudaOp", + "bytes": 264305664, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 71.16799801588058, + "candidate_gbps": 3713.8274416686877, + "candidate_peak_mib": 252.0615234375, + "provider_us": 110.23999750614166, + "provider_gbps": 2397.547804600368, + "provider_peak_mib": 252.09326171875, + "candidate_backward_us": 882.1919858455658, + "candidate_backward_gbps": 299.6010712415026, + "candidate_backward_peak_mib": 254.89501953125, + "provider_backward_us": 1276.960015296936, + "provider_backward_gbps": 206.98037591924137, + "provider_backward_peak_mib": 1.07568359375 + }, + { + "op": "adaln_row_gather", + "case": "S=32768", + "backend": "H3AdaLNRowGatherCudaOp", + "bytes": 2113929216, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 330.59200644493103, + "candidate_gbps": 6394.374863241383, + "candidate_peak_mib": 2016.0, + "provider_us": 621.0240125656128, + "provider_gbps": 3403.9411894345358, + "provider_peak_mib": 2016.25, + "candidate_backward_us": 1752.0639896392822, + "candidate_backward_gbps": 1206.5365354807727, + "candidate_backward_peak_mib": 2033.66845703125, + "provider_backward_us": 8044.352054595947, + "provider_backward_gbps": 262.7842741905182, + "provider_backward_peak_mib": 0.857421875 + }, + { + "op": "adaln_row_gather", + "case": "S=131072", + "backend": "H3AdaLNRowGatherCudaOp", + "bytes": 8455716864, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 1174.1759777069092, + "candidate_gbps": 7201.405091350512, + "candidate_peak_mib": 8064.0, + "provider_us": 2416.543960571289, + "provider_gbps": 3499.0949893586894, + "provider_peak_mib": 8065.0, + "candidate_backward_us": 5620.735883712769, + "candidate_backward_gbps": 1504.3789708216264, + "candidate_backward_peak_mib": 8130.05517578125, + "provider_backward_us": 32447.711944580078, + "provider_backward_gbps": 260.59516549093394, + "provider_backward_peak_mib": 0.5537109375 + } + ] +} diff --git a/reports/experiments/h3-final-adaln-out-b200/chain_replay.json b/reports/experiments/h3-final-adaln-out-b200/chain_replay.json new file mode 100644 index 000000000..5a34f1bd6 --- /dev/null +++ b/reports/experiments/h3-final-adaln-out-b200/chain_replay.json @@ -0,0 +1,5206 @@ +{ + "kind": "h3_conditioning_chain_replay", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "weights_sha256": { + "diffusion_pytorch_model-00001-of-00014.safetensors": { + "sha256": "2d847200c45c09dd7f973c1b096663068408ef851ee0b3711d059b6dc5dcd028", + "size_bytes": 4825958704 + } + }, + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "001684df3351e5184aad97949ce46090c841f78c", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "stages": [ + "timestep_sinusoid_h3", + "timestep_mlp_fp32", + "adaln_projection_3mod", + "adaln_row_gather", + "h3_rmsnorm", + "adaln_gate_residual", + "final_adaln_out" + ], + "cases": [ + { + "num_timesteps": 1, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.02018570899963379, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.02018570899963379, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 1.52587890625e-05, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0301055908203125, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0301055908203125, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07610607147216797, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07610607147216797, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 3.0517578125e-05, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.16942214965820312, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.16942214965820312, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01953125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09281444549560547, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09281444549560547, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 3.0517578125e-05, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2684783935546875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2684783935546875, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 1, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 7.301568984985352e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 3.91155481338501e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.004549264907836914, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01953125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05544424057006836, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05544424057006836, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09598255157470703, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09598255157470703, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 3.0517578125e-05, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0048828125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.022160768508911133, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.022160768508911133, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0048828125, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.04535102844238281, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.04535102844238281, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07610607147216797, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07610607147216797, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.16942214965820312, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.16942214965820312, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01953125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08827972412109375, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08827972412109375, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.009765625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2271709442138672, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2271709442138672, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 2, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007626533508300781, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0234375, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.06358528137207031, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.06358528137207031, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09641456604003906, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09641456604003906, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01171875, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003897547721862793, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003897547721862793, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05389690399169922, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05389690399169922, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0048828125, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.034628868103027344, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.034628868103027344, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0699777603149414, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0699777603149414, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0068359375, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.14091110229492188, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.14091110229492188, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05389690399169922, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05389690399169922, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09075355529785156, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09075355529785156, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.009765625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2684783935546875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2684783935546875, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 3, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07086944580078125, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07086944580078125, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09542465209960938, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.09542465209960938, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01171875, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.27352142333984375, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.27352142333984375, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 3, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.00390625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.003903031349182129, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0048828125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.022160768508911133, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.022160768508911133, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.040612220764160156, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0048828125, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.08124351501464844, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 257, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.04535102844238281, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.04535102844238281, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07404041290283203, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07404041290283203, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0078125, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.18407440185546875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.18407440185546875, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 4097, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.05122566223144531, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.11510276794433594, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.11510276794433594, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.009765625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2503814697265625, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2503814697265625, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + }, + { + "num_timesteps": 4, + "seq_len": 32768, + "seed": 0, + "stages": [ + { + "stage": "timestep_sinusoid_h3", + "backend": "H3TimestepSinusoidCudaOp", + "kernel_id": "rl_engine._C.h3_timestep_sinusoid_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.960464477539063e-08, + "within_tolerance": true + } + }, + { + "stage": "timestep_mlp_fp32", + "backend": "H3TimestepMLPCudaOp", + "kernel_id": "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 6.109476089477539e-07, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 1.1213123798370361e-06, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 5.103647708892822e-07, + "within_tolerance": true + } + }, + { + "stage": "adaln_projection_3mod", + "backend": "H3AdaLNProjectionCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "adaln_row_gather", + "backend": "H3AdaLNRowGatherCudaOp", + "kernel_id": "rl_engine._C.h3_adaln_row_gather_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.0078105926513671875, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.007814407348632812, + "within_tolerance": true + } + }, + { + "stage": "h3_rmsnorm", + "backend": "H3RMSNormCudaOp", + "kernel_id": "rl_engine._C.h3_rmsnorm_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.03125, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07958316802978516, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.07958316802978516, + "within_tolerance": true + } + }, + { + "stage": "adaln_gate_residual", + "backend": "H3GateResidualCudaOp", + "kernel_id": "rl_engine._C.h3_gate_residual_forward", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.015625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": true, + "max_abs": 0.0, + "within_tolerance": true + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.11510276794433594, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.11510276794433594, + "within_tolerance": true + } + }, + { + "stage": "final_adaln_out", + "backend": "H3FinalAdaLNOutCudaOp", + "kernel_id": "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward+rl_engine._C.h3_rmsnorm_forward[modulated]", + "repeat_bitwise_equal": true, + "chained_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.0625, + "within_tolerance": false + }, + "isolated_vs_provider": { + "bitwise_equal": false, + "max_abs": 0.01171875, + "within_tolerance": false + }, + "chained_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + }, + "provider_vs_golden": { + "bitwise_equal": false, + "max_abs": 0.2807731628417969, + "within_tolerance": true + } + } + ], + "first_drift": "timestep_mlp_fp32", + "first_isolated_drift": "timestep_mlp_fp32" + } + ], + "backward_cases": [ + { + "num_timesteps": 1, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 14.676741600036621, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.5142298936843872 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.5142298936843872 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.020853519439697266, + "max_abs_vs_golden_over_absmax": 0.0014208548236638308, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 14.676741600036621, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.02845982275903225 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.193450927734375e-05, + "max_abs_vs_golden_over_absmax": 1.494508069299627e-06, + "correctly_rounded_fraction": 0.02845982275903225 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.020853519439697266, + "max_abs_vs_golden_over_absmax": 0.0014208548236638308, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 20.589889526367188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.147125244140625e-05, + "max_abs_vs_golden_over_absmax": 1.5284808796423022e-06, + "correctly_rounded_fraction": 0.016372647136449814 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.147125244140625e-05, + "max_abs_vs_golden_over_absmax": 1.5284808796423022e-06, + "correctly_rounded_fraction": 0.016372647136449814 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.066619873046875, + "max_abs_vs_golden_over_absmax": 0.0032355624716728926, + "correctly_rounded_fraction": 2.4081899027805775e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 73.953369140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000110626220703125, + "max_abs_vs_golden_over_absmax": 1.4958915244278614e-06, + "correctly_rounded_fraction": 0.0524553582072258 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000110626220703125, + "max_abs_vs_golden_over_absmax": 1.4958915244278614e-06, + "correctly_rounded_fraction": 0.0524553582072258 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.23928070068359375, + "max_abs_vs_golden_over_absmax": 0.0032355617731809616, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 1.140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 5.351027357392013e-05, + "correctly_rounded_fraction": 0.9957404732704163 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 5.351027357392013e-05, + "correctly_rounded_fraction": 0.9957404732704163 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 2.6755136786960065e-05, + "correctly_rounded_fraction": 0.998334527015686 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.5096726417541504 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5346245765686035, + "max_abs_vs_golden_over_absmax": 0.004658695310354233, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 114.75843048095703, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.151611328125, + "max_abs_vs_golden_over_absmax": 0.0013211345067247748, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.392333984375e-05, + "max_abs_vs_golden_over_absmax": 7.313043397516594e-07, + "correctly_rounded_fraction": 0.0193452388048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5346245765686035, + "max_abs_vs_golden_over_absmax": 0.004658695310354233, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 174.07337951660156, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.22721099853515625, + "max_abs_vs_golden_over_absmax": 0.0013052598806098104, + "correctly_rounded_fraction": 1.1487342817417812e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00023651123046875, + "max_abs_vs_golden_over_absmax": 1.358686972707801e-06, + "correctly_rounded_fraction": 0.013279437087476254 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 2.0242090225219727, + "max_abs_vs_golden_over_absmax": 0.011628481559455395, + "correctly_rounded_fraction": 6.228077040759672e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 625.2249755859375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8160858154296875, + "max_abs_vs_golden_over_absmax": 0.001305267447605729, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.3666955283042626e-06, + "correctly_rounded_fraction": 0.0372023805975914 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 7.270416259765625, + "max_abs_vs_golden_over_absmax": 0.01162848062813282, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 11.4375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.005464480724185705, + "correctly_rounded_fraction": 0.7395860552787781 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 4.269125565770082e-05, + "correctly_rounded_fraction": 0.9957284331321716 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.28125, + "max_abs_vs_golden_over_absmax": 0.02459016442298889, + "correctly_rounded_fraction": 0.1703302264213562 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.125, + "max_abs_vs_golden_over_absmax": 0.025280898436903954, + "correctly_rounded_fraction": 0.17185433208942413 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.5178571343421936 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 17.811981201171875, + "max_abs_vs_golden_over_absmax": 0.033383384346961975, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 533.5582885742188, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8239898681640625, + "max_abs_vs_golden_over_absmax": 0.001544329570606351, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000701904296875, + "max_abs_vs_golden_over_absmax": 1.3155156466382323e-06, + "correctly_rounded_fraction": 0.0357142873108387 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 17.811981201171875, + "max_abs_vs_golden_over_absmax": 0.033383384346961975, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 627.8563232421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.204010009765625, + "max_abs_vs_golden_over_absmax": 0.0019176520872861147, + "correctly_rounded_fraction": 2.560431767051341e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0008544921875, + "max_abs_vs_golden_over_absmax": 1.360967758046172e-06, + "correctly_rounded_fraction": 0.015629082918167114 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 22.429458618164062, + "max_abs_vs_golden_over_absmax": 0.03572387248277664, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2255.091796875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.324493408203125, + "max_abs_vs_golden_over_absmax": 0.001917657325975597, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.3532743423638749e-06, + "correctly_rounded_fraction": 0.0520833358168602 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 80.56060791015625, + "max_abs_vs_golden_over_absmax": 0.03572387248277664, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 39.25, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.006369426846504211, + "correctly_rounded_fraction": 0.7380850911140442 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 4.976114723831415e-05, + "correctly_rounded_fraction": 0.9957199096679688 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.75, + "max_abs_vs_golden_over_absmax": 0.09554140269756317, + "correctly_rounded_fraction": 0.04575129225850105 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001953125, + "max_abs_vs_golden_over_absmax": 1.2765523024427239e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.5, + "max_abs_vs_golden_over_absmax": 0.09477124363183975, + "correctly_rounded_fraction": 0.0458829365670681 + } + } + } + }, + { + "num_timesteps": 1, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 1096.3087158203125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.691497802734375, + "max_abs_vs_golden_over_absmax": 0.00154290278442204, + "correctly_rounded_fraction": 0.5 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0022430419921875, + "max_abs_vs_golden_over_absmax": 2.045994961008546e-06, + "correctly_rounded_fraction": 0.5103236436843872 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 111.833740234375, + "max_abs_vs_golden_over_absmax": 0.1020093485713005, + "correctly_rounded_fraction": 0.5 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 1096.3087158203125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.691497802734375, + "max_abs_vs_golden_over_absmax": 0.00154290278442204, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0022430419921875, + "max_abs_vs_golden_over_absmax": 2.045994961008546e-06, + "correctly_rounded_fraction": 0.0206473208963871 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 111.833740234375, + "max_abs_vs_golden_over_absmax": 0.1020093485713005, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2459.76708984375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.49169921875, + "max_abs_vs_golden_over_absmax": 0.0010129817528650165, + "correctly_rounded_fraction": 5.453027551993728e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0032958984375, + "max_abs_vs_golden_over_absmax": 1.3399229601418483e-06, + "correctly_rounded_fraction": 0.016593534499406815 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 259.722412109375, + "max_abs_vs_golden_over_absmax": 0.10558821260929108, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8834.82421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.94921875, + "max_abs_vs_golden_over_absmax": 0.0010129481088370085, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.01171875, + "max_abs_vs_golden_over_absmax": 1.3264270819490775e-06, + "correctly_rounded_fraction": 0.063988097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 932.8533325195312, + "max_abs_vs_golden_over_absmax": 0.10558822005987167, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 121.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004115226212888956, + "correctly_rounded_fraction": 0.737601637840271 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.0005144032766111195, + "correctly_rounded_fraction": 0.9956783652305603 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 28.25, + "max_abs_vs_golden_over_absmax": 0.2325102835893631, + "correctly_rounded_fraction": 0.015966668725013733 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0078125, + "max_abs_vs_golden_over_absmax": 1.641281596675981e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0078125, + "max_abs_vs_golden_over_absmax": 1.641281596675981e-05, + "correctly_rounded_fraction": 0.999968945980072 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 110.0, + "max_abs_vs_golden_over_absmax": 0.2310924381017685, + "correctly_rounded_fraction": 0.015573329292237759 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024131298065185547, + "max_abs_vs_golden_over_absmax": 0.0027526128105819225, + "correctly_rounded_fraction": 2.1798271063744323e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024131298065185547, + "max_abs_vs_golden_over_absmax": 0.0027526128105819225, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 21.970394134521484, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04718971252441406, + "max_abs_vs_golden_over_absmax": 0.0021478773560374975, + "correctly_rounded_fraction": 1.972224526980426e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 78.91429901123047, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.16949462890625, + "max_abs_vs_golden_over_absmax": 0.0021478317212313414, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 3.578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 8.528930266038515e-06, + "correctly_rounded_fraction": 0.998363733291626 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 95.00251007080078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.126617431640625, + "max_abs_vs_golden_over_absmax": 0.0013327798806130886, + "correctly_rounded_fraction": 5.158923886483535e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00016021728515625, + "max_abs_vs_golden_over_absmax": 1.6864531744431588e-06, + "correctly_rounded_fraction": 0.02457391656935215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.8007991313934326, + "max_abs_vs_golden_over_absmax": 0.008429241366684437, + "correctly_rounded_fraction": 7.2660900514165405e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 95.00251007080078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.12661361694335938, + "max_abs_vs_golden_over_absmax": 0.0013327397173270583, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00012969970703125, + "max_abs_vs_golden_over_absmax": 1.365223965876794e-06, + "correctly_rounded_fraction": 0.02752976305782795 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.8007990717887878, + "max_abs_vs_golden_over_absmax": 0.008429241366684437, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.80873107910156, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2622566223144531, + "max_abs_vs_golden_over_absmax": 0.001508880639448762, + "correctly_rounded_fraction": 2.3874295948189683e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00018310546875, + "max_abs_vs_golden_over_absmax": 1.0534882903812104e-06, + "correctly_rounded_fraction": 0.014370596036314964 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.3165359497070312, + "max_abs_vs_golden_over_absmax": 0.007574624847620726, + "correctly_rounded_fraction": 2.4912308163038688e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.3004150390625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.941986083984375, + "max_abs_vs_golden_over_absmax": 0.0015088666696101427, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00067138671875, + "max_abs_vs_golden_over_absmax": 1.0754224604170304e-06, + "correctly_rounded_fraction": 0.0457589291036129 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 4.7288360595703125, + "max_abs_vs_golden_over_absmax": 0.007574616465717554, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 25.75, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.004854368977248669, + "correctly_rounded_fraction": 0.6666338443756104 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.8962378817377612e-05, + "correctly_rounded_fraction": 0.9989553093910217 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.019417475908994675, + "correctly_rounded_fraction": 0.22603164613246918 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.6273354887962341 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.01123595517128706, + "correctly_rounded_fraction": 0.2218191921710968 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 527.6658935546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.27825927734375, + "max_abs_vs_golden_over_absmax": 0.0005273398710414767, + "correctly_rounded_fraction": 4.359654212748865e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00042724609375, + "max_abs_vs_golden_over_absmax": 8.096905617094308e-07, + "correctly_rounded_fraction": 0.0324409119784832 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 9.172073364257812, + "max_abs_vs_golden_over_absmax": 0.01738234981894493, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 527.6620483398438, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2781982421875, + "max_abs_vs_golden_over_absmax": 0.000527228054124862, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00042724609375, + "max_abs_vs_golden_over_absmax": 8.096964734249923e-07, + "correctly_rounded_fraction": 0.0355282761156559 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 9.17205810546875, + "max_abs_vs_golden_over_absmax": 0.01738244853913784, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 630.265869140625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8614501953125, + "max_abs_vs_golden_over_absmax": 0.001366804470308125, + "correctly_rounded_fraction": 2.255947947560344e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000732421875, + "max_abs_vs_golden_over_absmax": 1.162084004135977e-06, + "correctly_rounded_fraction": 0.01656709983944893 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 21.645339965820312, + "max_abs_vs_golden_over_absmax": 0.03434319049119949, + "correctly_rounded_fraction": 8.304102721012896e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2263.81787109375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.09423828125, + "max_abs_vs_golden_over_absmax": 0.0013668229803442955, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0025634765625, + "max_abs_vs_golden_over_absmax": 1.1323687658659765e-06, + "correctly_rounded_fraction": 0.0699404776096344 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 77.7467041015625, + "max_abs_vs_golden_over_absmax": 0.0343431793153286, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 115.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004347825888544321, + "correctly_rounded_fraction": 0.6659318208694458 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.03125, + "max_abs_vs_golden_over_absmax": 0.00027173911803402007, + "correctly_rounded_fraction": 0.9989242553710938 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.75, + "max_abs_vs_golden_over_absmax": 0.05000000074505806, + "correctly_rounded_fraction": 0.0626961812376976 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6259506940841675 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 9.0, + "max_abs_vs_golden_over_absmax": 0.05882352963089943, + "correctly_rounded_fraction": 0.06214864179491997 + } + } + } + }, + { + "num_timesteps": 2, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 744.156982421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.7833938598632812, + "max_abs_vs_golden_over_absmax": 0.0023965290747582912, + "correctly_rounded_fraction": 1.2352353223832324e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001129150390625, + "max_abs_vs_golden_over_absmax": 1.5173551446423517e-06, + "correctly_rounded_fraction": 0.0301063172519207 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 101.708251953125, + "max_abs_vs_golden_over_absmax": 0.13667580485343933, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 744.156982421875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.7833938598632812, + "max_abs_vs_golden_over_absmax": 0.0023965290747582912, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00109100341796875, + "max_abs_vs_golden_over_absmax": 1.4660930673926487e-06, + "correctly_rounded_fraction": 0.0329241082072258 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 98.1940689086914, + "max_abs_vs_golden_over_absmax": 0.1319534331560135, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2462.752685546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.5404052734375, + "max_abs_vs_golden_over_absmax": 0.001437580562196672, + "correctly_rounded_fraction": 2.7334339392837137e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.002685546875, + "max_abs_vs_golden_over_absmax": 1.0904655027843546e-06, + "correctly_rounded_fraction": 0.0176222063601017 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 200.07321166992188, + "max_abs_vs_golden_over_absmax": 0.08123967051506042, + "correctly_rounded_fraction": 3.460042989900103e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8845.9306640625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 12.716552734375, + "max_abs_vs_golden_over_absmax": 0.0014375596074387431, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.009765625, + "max_abs_vs_golden_over_absmax": 1.103968088500551e-06, + "correctly_rounded_fraction": 0.0658482164144516 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 718.626708984375, + "max_abs_vs_golden_over_absmax": 0.08123811334371567, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 298.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.003355704713612795, + "correctly_rounded_fraction": 0.6657332181930542 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0016778523568063974, + "correctly_rounded_fraction": 0.9988868236541748 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 50.0, + "max_abs_vs_golden_over_absmax": 0.16778524219989777, + "correctly_rounded_fraction": 0.02199636586010456 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.6236565709114075 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0010504202218726277, + "correctly_rounded_fraction": 0.9999483227729797 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 60.75, + "max_abs_vs_golden_over_absmax": 0.12762604653835297, + "correctly_rounded_fraction": 0.020791996270418167 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 11.570573806762695, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.2099742889404297e-05, + "max_abs_vs_golden_over_absmax": 1.0457340522407321e-06, + "correctly_rounded_fraction": 0.0445265993475914 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.2099742889404297e-05, + "max_abs_vs_golden_over_absmax": 1.0457340522407321e-06, + "correctly_rounded_fraction": 0.0445265993475914 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.025829792022705078, + "max_abs_vs_golden_over_absmax": 0.0022323690354824066, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 11.57003402709961, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.208484172821045e-05, + "max_abs_vs_golden_over_absmax": 1.0444948657095665e-06, + "correctly_rounded_fraction": 0.0403645858168602 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.208484172821045e-05, + "max_abs_vs_golden_over_absmax": 1.0444948657095665e-06, + "correctly_rounded_fraction": 0.0403645858168602 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.025829792022705078, + "max_abs_vs_golden_over_absmax": 0.0022324733436107635, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 20.580326080322266, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.52587890625e-05, + "max_abs_vs_golden_over_absmax": 7.414259926008526e-07, + "correctly_rounded_fraction": 0.019046567380428314 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.52587890625e-05, + "max_abs_vs_golden_over_absmax": 7.414259926008526e-07, + "correctly_rounded_fraction": 0.019046567380428314 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04852867126464844, + "max_abs_vs_golden_over_absmax": 0.002358012832701206, + "correctly_rounded_fraction": 2.159066752938088e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 73.91921997070312, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.7220458984375e-05, + "max_abs_vs_golden_over_absmax": 7.740944738543476e-07, + "correctly_rounded_fraction": 0.1011904776096344 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.7220458984375e-05, + "max_abs_vs_golden_over_absmax": 7.740944738543476e-07, + "correctly_rounded_fraction": 0.1011904776096344 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1743011474609375, + "max_abs_vs_golden_over_absmax": 0.002357994904741645, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 2.859375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 2.134562782885041e-05, + "correctly_rounded_fraction": 0.9968344569206238 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 2.134562782885041e-05, + "correctly_rounded_fraction": 0.9968344569206238 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 1.0672813914425205e-05, + "correctly_rounded_fraction": 0.9970113635063171 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087639331817627, + "max_abs_vs_golden_over_absmax": 0.0014897150686010718, + "correctly_rounded_fraction": 4.359654212748865e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00018215179443359375, + "max_abs_vs_golden_over_absmax": 2.494892214599531e-06, + "correctly_rounded_fraction": 0.02586873434484005 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.36901092529296875, + "max_abs_vs_golden_over_absmax": 0.0050542596727609634, + "correctly_rounded_fraction": 4.359654212748865e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 73.0098876953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.1087636947631836, + "max_abs_vs_golden_over_absmax": 0.0014897118089720607, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0001819133758544922, + "max_abs_vs_golden_over_absmax": 2.4916266738728154e-06, + "correctly_rounded_fraction": 0.0232514888048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.3690032958984375, + "max_abs_vs_golden_over_absmax": 0.005054154898971319, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.8400421142578, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2456045150756836, + "max_abs_vs_golden_over_absmax": 0.001412819023244083, + "correctly_rounded_fraction": 2.089865847665351e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000171661376953125, + "max_abs_vs_golden_over_absmax": 9.874673878584872e-07, + "correctly_rounded_fraction": 0.012266821227967739 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.1511993408203125, + "max_abs_vs_golden_over_absmax": 0.00662217615172267, + "correctly_rounded_fraction": 3.1140386909100926e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.4014892578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.8821792602539062, + "max_abs_vs_golden_over_absmax": 0.0014128397451713681, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0006103515625, + "max_abs_vs_golden_over_absmax": 9.774985301191919e-07, + "correctly_rounded_fraction": 0.0412946417927742 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 4.1348876953125, + "max_abs_vs_golden_over_absmax": 0.006622161716222763, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 27.375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.004566209856420755, + "correctly_rounded_fraction": 0.6352493762969971 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.7836757251643576e-05, + "correctly_rounded_fraction": 0.9988961815834045 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.3125, + "max_abs_vs_golden_over_absmax": 0.01141552533954382, + "correctly_rounded_fraction": 0.2612820565700531 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.6106977462768555 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.01123595517128706, + "correctly_rounded_fraction": 0.25729578733444214 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 216.79254150390625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.34185791015625, + "max_abs_vs_golden_over_absmax": 0.0015768896555528045, + "correctly_rounded_fraction": 4.650297705666162e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.293441634217743e-06, + "correctly_rounded_fraction": 0.02635483630001545 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 7.005500793457031, + "max_abs_vs_golden_over_absmax": 0.03231430798768997, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 189.4628448486328, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.3418426513671875, + "max_abs_vs_golden_over_absmax": 0.0018042727606371045, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009307861328125, + "max_abs_vs_golden_over_absmax": 4.912763415632071e-06, + "correctly_rounded_fraction": 0.0204613097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 6.1077117919921875, + "max_abs_vs_golden_over_absmax": 0.032236989587545395, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 632.3684692382812, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.124969482421875, + "max_abs_vs_golden_over_absmax": 0.0017789778066799045, + "correctly_rounded_fraction": 2.3805096134310588e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0009765625, + "max_abs_vs_golden_over_absmax": 1.544293468214164e-06, + "correctly_rounded_fraction": 0.011527756229043007 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 15.074493408203125, + "max_abs_vs_golden_over_absmax": 0.023838147521018982, + "correctly_rounded_fraction": 2.2144274680613307e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2271.419677734375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 4.0406494140625, + "max_abs_vs_golden_over_absmax": 0.0017789092380553484, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0037841796875, + "max_abs_vs_golden_over_absmax": 1.6659976154187461e-06, + "correctly_rounded_fraction": 0.0338541679084301 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 54.1448974609375, + "max_abs_vs_golden_over_absmax": 0.023837469518184662, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 94.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.005291005130857229, + "correctly_rounded_fraction": 0.6358715295791626 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 4.13359775848221e-05, + "correctly_rounded_fraction": 0.9988698363304138 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 4.5, + "max_abs_vs_golden_over_absmax": 0.0476190485060215, + "correctly_rounded_fraction": 0.07537095248699188 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6086929440498352 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 2.5531046048854478e-05, + "correctly_rounded_fraction": 0.9999896287918091 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 6.375, + "max_abs_vs_golden_over_absmax": 0.0416666679084301, + "correctly_rounded_fraction": 0.0746527761220932 + } + } + } + }, + { + "num_timesteps": 3, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 1221.45751953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.3370361328125, + "max_abs_vs_golden_over_absmax": 0.0010946234688162804, + "correctly_rounded_fraction": 6.5394810917496216e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001739501953125, + "max_abs_vs_golden_over_absmax": 1.424119886905828e-06, + "correctly_rounded_fraction": 0.0268104188144207 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 51.452091217041016, + "max_abs_vs_golden_over_absmax": 0.042123522609472275, + "correctly_rounded_fraction": 7.266090165103378e-07 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 1221.45751953125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.3369140625, + "max_abs_vs_golden_over_absmax": 0.0010945235844701529, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.001708984375, + "max_abs_vs_golden_over_absmax": 1.3991353853270994e-06, + "correctly_rounded_fraction": 0.032738097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 51.44636154174805, + "max_abs_vs_golden_over_absmax": 0.04211882874369621, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2460.399169921875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.100341796875, + "max_abs_vs_golden_over_absmax": 0.0012600970221683383, + "correctly_rounded_fraction": 2.9894770705141127e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00274658203125, + "max_abs_vs_golden_over_absmax": 1.1163156159454957e-06, + "correctly_rounded_fraction": 0.01790323108434677 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 141.7109375, + "max_abs_vs_golden_over_absmax": 0.05759672448039055, + "correctly_rounded_fraction": 2.7680343350766634e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8837.2431640625, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 11.13623046875, + "max_abs_vs_golden_over_absmax": 0.0012601475464180112, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00927734375, + "max_abs_vs_golden_over_absmax": 1.0498006304260343e-06, + "correctly_rounded_fraction": 0.0632440522313118 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 508.999755859375, + "max_abs_vs_golden_over_absmax": 0.057597119361162186, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 288.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.0069444444961845875, + "correctly_rounded_fraction": 0.6357954144477844 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.0017361111240461469, + "correctly_rounded_fraction": 0.9988410472869873 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 36.0, + "max_abs_vs_golden_over_absmax": 0.125, + "correctly_rounded_fraction": 0.02689298801124096 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.6103463768959045 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.00026260505546815693, + "correctly_rounded_fraction": 0.9999379515647888 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 59.0, + "max_abs_vs_golden_over_absmax": 0.1239495798945427, + "correctly_rounded_fraction": 0.02604166604578495 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 3, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.821487426757812e-06, + "max_abs_vs_golden_over_absmax": 1.0062508408736903e-06, + "correctly_rounded_fraction": 0.0394497849047184 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024135589599609375, + "max_abs_vs_golden_over_absmax": 0.002753102220594883, + "correctly_rounded_fraction": 1.4532180330206756e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 8.766688346862793, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 8.58306884765625e-06, + "max_abs_vs_golden_over_absmax": 9.790549029276008e-07, + "correctly_rounded_fraction": 0.0461309514939785 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.024135589599609375, + "max_abs_vs_golden_over_absmax": 0.002753102220594883, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 21.970394134521484, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.811981201171875e-05, + "max_abs_vs_golden_over_absmax": 8.24737696802913e-07, + "correctly_rounded_fraction": 0.019412847235798836 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.04718971252441406, + "max_abs_vs_golden_over_absmax": 0.0021478773560374975, + "correctly_rounded_fraction": 1.972224526980426e-05 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 78.91429901123047, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 8.21775699932914e-07, + "correctly_rounded_fraction": 0.110863097012043 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.16949462890625, + "max_abs_vs_golden_over_absmax": 0.0021478317212313414, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 3.578125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.103515625e-05, + "max_abs_vs_golden_over_absmax": 1.705786053207703e-05, + "correctly_rounded_fraction": 0.9982069134712219 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.0517578125e-05, + "max_abs_vs_golden_over_absmax": 8.528930266038515e-06, + "correctly_rounded_fraction": 0.998363733291626 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 4.46875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 257, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 54.46367263793945, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.11591720581054688, + "max_abs_vs_golden_over_absmax": 0.0021283398382365704, + "correctly_rounded_fraction": 3.6330450257082703e-06 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 6.4849853515625e-05, + "max_abs_vs_golden_over_absmax": 1.1906992085641832e-06, + "correctly_rounded_fraction": 0.02683221735060215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.410614013671875, + "max_abs_vs_golden_over_absmax": 0.007539227604866028, + "correctly_rounded_fraction": 2.1798271063744323e-06 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 54.46367263793945, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.11591720581054688, + "max_abs_vs_golden_over_absmax": 0.0021283398382365704, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 5.340576171875e-05, + "max_abs_vs_golden_over_absmax": 9.805758054426406e-07, + "correctly_rounded_fraction": 0.0271577388048172 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.4106101989746094, + "max_abs_vs_golden_over_absmax": 0.007539157290011644, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 173.8180694580078, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.2628202438354492, + "max_abs_vs_golden_over_absmax": 0.001512042130343616, + "correctly_rounded_fraction": 1.4670581549580675e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00014495849609375, + "max_abs_vs_golden_over_absmax": 8.339667942891538e-07, + "correctly_rounded_fraction": 0.012976337224245071 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 1.0045738220214844, + "max_abs_vs_golden_over_absmax": 0.005779455415904522, + "correctly_rounded_fraction": 6.920085979800206e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 624.3176879882812, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.9439964294433594, + "max_abs_vs_golden_over_absmax": 0.0015120450407266617, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000537872314453125, + "max_abs_vs_golden_over_absmax": 8.615362503405777e-07, + "correctly_rounded_fraction": 0.0386904776096344 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.6082000732421875, + "max_abs_vs_golden_over_absmax": 0.005779429338872433, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 25.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.125, + "max_abs_vs_golden_over_absmax": 0.0049019609577953815, + "correctly_rounded_fraction": 0.619565486907959 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00048828125, + "max_abs_vs_golden_over_absmax": 1.914828499138821e-05, + "correctly_rounded_fraction": 0.9988008141517639 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.375, + "max_abs_vs_golden_over_absmax": 0.014705882407724857, + "correctly_rounded_fraction": 0.2856287956237793 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 44.5, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.25, + "max_abs_vs_golden_over_absmax": 0.00561797758564353, + "correctly_rounded_fraction": 0.597811222076416 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0, + "max_abs_vs_golden_over_absmax": 0.0, + "correctly_rounded_fraction": 0.9999999403953552 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.01123595517128706, + "correctly_rounded_fraction": 0.28108465671539307 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 4097, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 326.7618713378906, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.541351318359375, + "max_abs_vs_golden_over_absmax": 0.0016567150596529245, + "correctly_rounded_fraction": 3.705705967149697e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000244140625, + "max_abs_vs_golden_over_absmax": 7.471514891221887e-07, + "correctly_rounded_fraction": 0.0328994020819664 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.7530059814453125, + "max_abs_vs_golden_over_absmax": 0.011485446244478226, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 326.7618713378906, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5413436889648438, + "max_abs_vs_golden_over_absmax": 0.0016566917765885592, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.000213623046875, + "max_abs_vs_golden_over_absmax": 6.537575814036245e-07, + "correctly_rounded_fraction": 0.0319940485060215 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.7530059814453125, + "max_abs_vs_golden_over_absmax": 0.011485446244478226, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 630.1971435546875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.950373649597168, + "max_abs_vs_golden_over_absmax": 0.0015080576995387673, + "correctly_rounded_fraction": 2.4981509341159835e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0006103515625, + "max_abs_vs_golden_over_absmax": 9.685089708000305e-07, + "correctly_rounded_fraction": 0.013629386201500893 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 14.008440971374512, + "max_abs_vs_golden_over_absmax": 0.02222866378724575, + "correctly_rounded_fraction": 2.4912308163038688e-06 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 2263.544677734375, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 3.4134902954101562, + "max_abs_vs_golden_over_absmax": 0.0015080287121236324, + "correctly_rounded_fraction": 0.0 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.002197265625, + "max_abs_vs_golden_over_absmax": 9.70718929238501e-07, + "correctly_rounded_fraction": 0.0498511902987957 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 50.315120697021484, + "max_abs_vs_golden_over_absmax": 0.022228464484214783, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 104.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.5, + "max_abs_vs_golden_over_absmax": 0.004807692486792803, + "correctly_rounded_fraction": 0.6189784407615662 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.03125, + "max_abs_vs_golden_over_absmax": 0.0003004807804245502, + "correctly_rounded_fraction": 0.9987642765045166 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 3.0, + "max_abs_vs_golden_over_absmax": 0.028846153989434242, + "correctly_rounded_fraction": 0.0858362540602684 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 153.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.006535947788506746, + "correctly_rounded_fraction": 0.6000640392303467 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 2.5531046048854478e-05, + "correctly_rounded_fraction": 0.9999793171882629 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 5.0, + "max_abs_vs_golden_over_absmax": 0.03267974033951759, + "correctly_rounded_fraction": 0.08580315858125687 + } + } + } + }, + { + "num_timesteps": 4, + "seq_len": 32768, + "leaves": { + "time_embedder.linear_1.weight": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436309814453125, + "max_abs_vs_golden_over_absmax": 0.0007790509844198823, + "correctly_rounded_fraction": 3.4150623832829297e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.02686346136033535 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 35.3507080078125, + "max_abs_vs_golden_over_absmax": 0.04278850182890892, + "correctly_rounded_fraction": 7.266090165103378e-07 + } + }, + "time_embedder.linear_1.bias": { + "golden_absmax": 826.173095703125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.6436004638671875, + "max_abs_vs_golden_over_absmax": 0.0007790140807628632, + "correctly_rounded_fraction": 0.00018601190822664648 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0011138916015625, + "max_abs_vs_golden_over_absmax": 1.348254500044277e-06, + "correctly_rounded_fraction": 0.0282738097012043 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 35.3507080078125, + "max_abs_vs_golden_over_absmax": 0.04278850182890892, + "correctly_rounded_fraction": 0.0 + } + }, + "time_embedder.linear_2.weight": { + "golden_absmax": 2460.831298828125, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.702392578125, + "max_abs_vs_golden_over_absmax": 0.0010981624945998192, + "correctly_rounded_fraction": 3.480803206912242e-05 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0030517578125, + "max_abs_vs_golden_over_absmax": 1.2401328604028095e-06, + "correctly_rounded_fraction": 0.018226468935608864 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 129.7418212890625, + "max_abs_vs_golden_over_absmax": 0.052722759544849396, + "correctly_rounded_fraction": 2.7680343350766634e-07 + } + }, + "time_embedder.linear_2.bias": { + "golden_absmax": 8838.8046875, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 9.70703125, + "max_abs_vs_golden_over_absmax": 0.0010982289677485824, + "correctly_rounded_fraction": 0.00037202381645329297 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.01171875, + "max_abs_vs_golden_over_absmax": 1.3258297713036882e-06, + "correctly_rounded_fraction": 0.0569196455180645 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 465.998046875, + "max_abs_vs_golden_over_absmax": 0.052721839398145676, + "correctly_rounded_fraction": 0.0 + } + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "golden_absmax": 282.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 1.0, + "max_abs_vs_golden_over_absmax": 0.003546099178493023, + "correctly_rounded_fraction": 0.6200082898139954 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.0625, + "max_abs_vs_golden_over_absmax": 0.00022163119865581393, + "correctly_rounded_fraction": 0.9987457394599915 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 25.25, + "max_abs_vs_golden_over_absmax": 0.08953900635242462, + "correctly_rounded_fraction": 0.030449382960796356 + } + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "golden_absmax": 476.0, + "candidate": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 2.0, + "max_abs_vs_golden_over_absmax": 0.004201680887490511, + "correctly_rounded_fraction": 0.5982245802879333 + }, + "candidate_fused": { + "repeat_bitwise_equal": true, + "max_abs_vs_golden": 0.00390625, + "max_abs_vs_golden_over_absmax": 8.206407983379904e-06, + "correctly_rounded_fraction": 0.9999586343765259 + }, + "provider": { + "repeat_bitwise_equal": false, + "max_abs_vs_golden": 40.0, + "max_abs_vs_golden_over_absmax": 0.08403361588716507, + "correctly_rounded_fraction": 0.0304129458963871 + } + } + } + } + ] +} diff --git a/reports/experiments/h3-final-adaln-out-b200/figure.png b/reports/experiments/h3-final-adaln-out-b200/figure.png new file mode 100644 index 000000000..0e5a33544 Binary files /dev/null and b/reports/experiments/h3-final-adaln-out-b200/figure.png differ diff --git a/reports/experiments/h3-final-adaln-out-b200/report.json b/reports/experiments/h3-final-adaln-out-b200/report.json new file mode 100644 index 000000000..5487c5052 --- /dev/null +++ b/reports/experiments/h3-final-adaln-out-b200/report.json @@ -0,0 +1,6189 @@ +{ + "kind": "h3_operator_report", + "op": "final_adaln_out", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "weight_source": "pinned_checkpoint", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "bde1d8d97975a9b5ad32b7f0f8486e7f6815638b", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "equal_to_diffusers_fraction": 0.9996244311332703, + "max_abs_vs_golden": 0.4453125, + "provider_max_abs_vs_golden": 0.4453125, + "rows_batch_invariant": true, + "backward": { + "cuda": { + "repeat_bitwise_equal": true, + "rel_error": { + "dx": 0.002330429519732287, + "d_norm_w": 0.0022793000129727424, + "d_temb": 0.0009756835628098683, + "dW": 0.00310445472056345, + "db": 0.002154037178338269 + } + }, + "provider": { + "repeat_bitwise_equal": false, + "rel_error": { + "dx": 0.004053890312809679, + "d_norm_w": 0.002528636418480549, + "d_temb": 0.036257690920956674, + "dW": 0.06475266284250622, + "db": 0.05896219040633434 + } + } + } + }, + "perf": [ + { + "op": "final_adaln_out", + "case": "S=4097", + "backend": "H3FinalAdaLNOutCudaOp", + "bytes": 145904640, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 2455.615997314453, + 2463.42396736145, + 2087.3920917510986, + 2463.42396736145, + 2461.7600440979004, + 2465.1200771331787, + 2464.384078979492, + 208.00000429153442, + 151.2639969587326, + 182.559996843338, + 148.92800152301788, + 218.01599860191345, + 155.32800555229187, + 198.5280066728592, + 147.61599898338318, + 215.03999829292297, + 147.5840061903, + 200.76799392700195, + 146.62399888038635, + 201.75999402999878, + 144.896000623703, + 196.79999351501465, + 165.8879965543747, + 197.4720060825348, + 147.93600142002106, + 201.79200172424316, + 166.1120057106018, + 184.51200425624847, + 145.34400403499603, + 208.28799903392792, + 166.72000288963318, + 204.92799580097198, + 147.5840061903, + 226.6560047864914, + 149.08799529075623, + 196.83200120925903, + 156.63999319076538, + 199.71199333667755, + 153.21600437164307, + 198.97599518299103, + 125.72799623012543, + 208.76799523830414, + 146.7839926481247, + 187.391996383667, + 127.10399925708771, + 200.6399929523468, + 149.9200016260147, + 201.4079988002777, + 149.82399344444275, + 197.11999595165253, + 160.5439931154251, + 209.08799767494202, + 153.50399911403656, + 200.3519982099533, + 130.2720010280609, + 203.16800475120544, + 159.4880074262619, + 203.48800718784332, + 148.8959938287735, + 198.43199849128723, + 164.99200463294983, + 204.96000349521637, + 144.896000623703, + 207.07200467586517, + 148.22399616241455, + 202.33599841594696, + 147.10399508476257, + 198.62399995326996, + 181.88799917697906, + 209.82399582862854, + 146.36799693107605, + 200.41599869728088, + 148.8640010356903, + 200.28799772262573, + 145.79200744628906, + 198.7839937210083, + 160.70400178432465, + 197.6960003376007, + 163.455992937088, + 215.61600267887115, + 132.9919993877411, + 202.91200280189514, + 145.53600549697876, + 197.02400267124176, + 146.91199362277985, + 201.1519968509674, + 170.1119989156723, + 221.11999988555908, + 149.59999918937683, + 216.09599888324738, + 147.2640037536621, + 204.12799715995789, + 170.3999936580658, + 396.1920142173767, + 148.3200043439865, + 201.12000405788422, + 148.41599762439728, + 199.96799528598785, + 151.45599842071533, + 202.30400562286377, + 148.8959938287735, + 204.0960043668747, + 148.80000054836273, + 179.967999458313, + 144.80000734329224, + 202.01599597930908, + 146.4959979057312, + 206.62400126457214, + 146.30399644374847, + 195.77600061893463, + 145.28000354766846, + 198.46400618553162, + 146.11199498176575, + 197.28000462055206, + 150.87999403476715, + 195.8719938993454, + 154.6880006790161, + 198.46400618553162, + 147.039994597435, + 202.78400182724, + 156.54399991035461, + 196.54400646686554, + 147.67999947071075, + 196.16000354290009, + 371.3279962539673, + 202.33599841594696, + 126.81600451469421, + 202.91200280189514, + 162.62400150299072, + 197.66399264335632, + 149.98400211334229, + 198.04799556732178, + 147.8080004453659, + 189.40800428390503, + 149.02399480342865, + 187.51999735832214, + 147.5840061903, + 199.0080028772354, + 145.91999351978302, + 205.05599677562714, + 150.56000649929047, + 202.30400562286377, + 147.20000326633453, + 196.57599925994873, + 150.04800260066986, + 204.25599813461304, + 154.40000593662262, + 201.53599977493286, + 188.03200125694275, + 198.4959989786148, + 155.07200360298157, + 216.3199931383133, + 167.7439957857132, + 195.96800208091736, + 145.21600306034088, + 220.0320065021515, + 149.4400054216385, + 200.03199577331543, + 146.17599546909332, + 207.74400234222412, + 147.93600142002106, + 198.33600521087646, + 181.95199966430664, + 209.6319943666458, + 149.08799529075623, + 201.9840031862259, + 148.12800288200378, + 208.80000293254852, + 159.743994474411, + 197.6960003376007, + 145.88800072669983, + 209.9200040102005, + 149.05600249767303, + 199.13600385189056, + 135.5839967727661, + 200.80000162124634, + 150.11200308799744, + 220.89600563049316, + 146.30399644374847, + 213.79199624061584, + 148.8959938287735, + 202.2079974412918, + 147.2959965467453, + 208.28799903392792, + 146.30399644374847, + 197.08800315856934, + 147.8080004453659, + 207.45599269866943, + 151.71200037002563, + 217.75999665260315, + 153.43999862670898, + 209.08799767494202, + 157.18400478363037, + 197.56799936294556, + 148.0959951877594, + 195.8719938993454, + 146.68799936771393, + 198.40000569820404, + 146.7839926481247, + 202.2079974412918 + ], + "provider": [ + 133.98399949073792, + 133.27999413013458, + 131.29599392414093, + 133.27999413013458, + 135.68000495433807, + 133.66399705410004, + 131.9040060043335, + 251.0719895362854, + 170.75200378894806, + 243.1039959192276, + 173.8560050725937, + 283.4239900112152, + 183.00800025463104, + 248.416006565094, + 169.63200271129608, + 252.76800990104675, + 168.32000017166138, + 259.0399980545044, + 166.24000668525696, + 256.0960054397583, + 172.4800020456314, + 244.54399943351746, + 173.8560050725937, + 257.05599784851074, + 167.71200299263, + 249.88800287246704, + 169.95200514793396, + 234.65600609779358, + 166.78400337696075, + 249.7279942035675, + 169.66399550437927, + 258.5600018501282, + 166.20799899101257, + 249.34400618076324, + 168.92799735069275, + 251.0719895362854, + 167.61599481105804, + 251.2960135936737, + 178.97599935531616, + 244.9599951505661, + 168.2880073785782, + 244.35199797153473, + 167.61599481105804, + 200.8640021085739, + 164.57599401474, + 244.80000138282776, + 173.95199835300446, + 255.8079957962036, + 166.75199568271637, + 228.5120040178299, + 169.76000368595123, + 250.68798661231995, + 171.07200622558594, + 250.62400102615356, + 165.3759926557541, + 248.79999458789825, + 166.4319932460785, + 253.91998887062073, + 166.55999422073364, + 246.75199389457703, + 167.39200055599213, + 244.09599602222443, + 190.5599981546402, + 246.5279996395111, + 174.94399845600128, + 256.22400641441345, + 165.8560037612915, + 253.9840042591095, + 168.89600455760956, + 256.28799200057983, + 166.24000668525696, + 250.36799907684326, + 172.86400496959686, + 256.3199996948242, + 168.47999393939972, + 253.02401185035706, + 170.33599317073822, + 244.03199553489685, + 176.1920005083084, + 249.66399371623993, + 168.47999393939972, + 271.0080146789551, + 179.26399409770966, + 245.7599937915802, + 169.855996966362, + 252.99200415611267, + 168.7999963760376, + 266.30398631095886, + 176.4480024576187, + 255.71200251579285, + 171.1679995059967, + 257.82400369644165, + 171.7119961977005, + 248.48000705242157, + 167.84000396728516, + 249.53599274158478, + 169.0559983253479, + 256.1280131340027, + 170.0800061225891, + 253.12000513076782, + 167.13599860668182, + 250.65600872039795, + 176.28799378871918, + 185.5040043592453, + 164.67200219631195, + 251.71199440956116, + 167.87199676036835, + 245.37600576877594, + 167.1680063009262, + 242.5920069217682, + 167.07199811935425, + 248.35200607776642, + 168.96000504493713, + 256.5760016441345, + 168.60799491405487, + 244.3840056657791, + 167.80799627304077, + 257.7280104160309, + 166.33599996566772, + 244.51200664043427, + 167.1999990940094, + 250.17601251602173, + 166.49599373340607, + 244.28799748420715, + 169.8240041732788, + 233.2479953765869, + 166.62399470806122, + 251.8399953842163, + 168.09600591659546, + 253.28001379966736, + 176.06399953365326, + 262.62399554252625, + 174.17599260807037, + 250.36799907684326, + 177.824005484581, + 245.60000002384186, + 173.567995429039, + 245.7599937915802, + 143.5839980840683, + 244.00000274181366, + 168.35199296474457, + 248.35200607776642, + 165.8560037612915, + 250.04801154136658, + 175.135999917984, + 255.90398907661438, + 166.55999422073364, + 248.25599789619446, + 167.04000532627106, + 243.16799640655518, + 169.0559983253479, + 254.88001108169556, + 168.92799735069275, + 243.23199689388275, + 167.1999990940094, + 244.60799992084503, + 168.32000017166138, + 248.51199984550476, + 180.41600286960602, + 245.31200528144836, + 180.35200238227844, + 246.65600061416626, + 171.4559942483902, + 246.65600061416626, + 168.38400065898895, + 250.17601251602173, + 163.35999965667725, + 246.39999866485596, + 167.67999529838562, + 244.9920028448105, + 165.24800658226013, + 244.6720004081726, + 166.75199568271637, + 239.6160066127777, + 168.38400065898895, + 247.6480007171631, + 169.63200271129608, + 245.5040067434311, + 186.88000738620758, + 248.35200607776642, + 167.71200299263, + 253.79198789596558, + 165.47200083732605, + 260.96001267433167, + 173.95199835300446, + 243.96799504756927, + 165.72800278663635, + 245.82399427890778, + 166.75199568271637, + 246.5600073337555, + 169.63200271129608, + 248.57600033283234, + 189.60000574588776, + 248.09600412845612, + 168.47999393939972, + 247.00799584388733, + 168.60799491405487, + 259.3599855899811, + 166.30400717258453, + 253.9519965648651 + ], + "candidate_backward": [ + 7954.559803009033, + 7563.424110412598, + 7577.727794647217, + 7951.488018035889, + 7946.335792541504, + 7582.816123962402, + 5440.544128417969, + 1073.2159614562988, + 1384.4480514526367, + 1426.0480403900146, + 1394.6559429168701, + 1453.05597782135, + 1099.552035331726, + 1202.3680210113525, + 1392.6719427108765, + 1380.511999130249, + 1362.8159761428833, + 1195.072054862976, + 1100.8000373840332, + 1203.1999826431274, + 1409.0240001678467, + 1440.0960206985474, + 1411.9999408721924, + 1379.2959451675415, + 1385.472059249878, + 4423.647880554199, + 1353.7280559539795, + 1416.159987449646, + 1395.6799507141113, + 1397.760033607483, + 1327.5200128555298, + 1202.0479440689087, + 1429.4719696044922, + 1453.0240297317505, + 1358.8160276412964, + 1417.248010635376, + 1434.4960451126099, + 1382.431983947754, + 1378.175973892212, + 1402.9120206832886, + 1395.616054534912, + 1419.2639589309692, + 1431.5840005874634, + 1426.3999462127686, + 1067.9680109024048, + 1411.1039638519287, + 1387.4880075454712, + 1398.6239433288574, + 1142.6559686660767, + 1455.1039934158325, + 1408.4479808807373, + 1350.0479459762573, + 1352.6719808578491, + 1398.7840414047241, + 1362.7519607543945, + 1399.839997291565, + 1349.6320247650146, + 1453.0880451202393, + 1357.7920198440552, + 1403.7760496139526, + 1430.5280447006226, + 1382.464051246643, + 1376.1919736862183, + 1465.343952178955, + 1376.255989074707, + 1426.4320135116577, + 1372.0959424972534, + 1367.9039478302002, + 1335.8399868011475, + 1445.4400539398193, + 1218.6239957809448, + 1456.0960531234741, + 1402.8799533843994, + 1384.4799995422363, + 1422.3040342330933, + 1457.1199417114258, + 1420.2879667282104, + 1449.9520063400269, + 1377.951979637146, + 1440.000057220459, + 1371.1040019989014, + 1199.1679668426514, + 1225.6959676742554, + 1457.6319456100464, + 1414.3999814987183, + 1439.7120475769043, + 1444.86403465271, + 1337.3440504074097, + 1392.6399946212769, + 1443.9040422439575, + 1315.9359693527222, + 1439.7120475769043, + 1337.280035018921, + 1471.4879989624023, + 1374.176025390625, + 1399.8080492019653, + 1394.495964050293, + 1415.295958518982, + 1396.7679738998413, + 2043.936014175415, + 1461.2480401992798, + 1369.1200017929077, + 1372.1599578857422, + 1410.048007965088, + 1392.7359580993652, + 1197.0560550689697, + 1455.1039934158325, + 1420.2560186386108, + 1341.439962387085, + 1407.9999923706055, + 1377.2159814834595, + 1442.1440362930298, + 1380.3520202636719, + 1391.6480541229248, + 1355.7440042495728, + 1430.5280447006226, + 1360.8640432357788, + 1410.207986831665, + 1358.8800430297852, + 1386.7839574813843, + 1379.1359663009644, + 1418.239951133728, + 1352.6400327682495, + 1452.0000219345093, + 1605.631947517395, + 1218.5920476913452, + 1342.3999547958374, + 1368.2880401611328, + 1369.055986404419, + 1386.5599632263184, + 1414.1119718551636, + 1385.4399919509888, + 1354.7519445419312, + 1419.2639589309692, + 1378.3040046691895, + 1447.9360580444336, + 1382.4000358581543, + 1359.8719835281372, + 1394.6880102157593, + 1445.855975151062, + 1384.28795337677, + 1427.8719425201416, + 1315.8719539642334, + 1249.2799758911133, + 1415.168046951294, + 1455.0399780273438, + 1408.992052078247, + 1453.0240297317505, + 1327.072024345398, + 1455.1039934158325, + 1405.9840440750122, + 1414.1440391540527, + 1357.8239679336548, + 1371.1680173873901, + 1413.1200313568115, + 1392.575979232788, + 1335.2960348129272, + 1367.0079708099365, + 1393.664002418518, + 1422.368049621582, + 1362.8480434417725, + 1396.5439796447754, + 1391.487956047058, + 1373.1520175933838, + 1349.727988243103, + 1408.031940460205, + 1437.7599954605103, + 1389.4399404525757, + 1368.6399459838867, + 1401.8880128860474, + 1392.6399946212769, + 1428.6079406738281, + 1427.4239540100098, + 1359.7760200500488, + 1395.6799507141113, + 1461.1200094223022, + 1344.4479703903198, + 1429.535984992981, + 1354.848027229309, + 1413.1200313568115, + 1416.159987449646, + 1475.6159782409668, + 1327.3279666900635, + 1403.9360284805298, + 1405.951976776123, + 1416.3520336151123, + 1339.359998703003, + 1433.3120584487915, + 1391.327977180481, + 1387.55202293396, + 1362.9440069198608, + 1402.400016784668, + 1348.6080169677734, + 1383.4240436553955, + 1401.8239974975586, + 1380.3839683532715, + 1342.3999547958374, + 1374.2079734802246, + 1396.7360258102417, + 1397.760033607483 + ], + "provider_backward": [ + 627.7120113372803, + 629.6640038490295, + 633.7599754333496, + 631.7440271377563, + 627.6159882545471, + 629.7280192375183, + 631.8079829216003, + 633.8559985160828, + 884.768009185791, + 846.8160033226013, + 896.0319757461548, + 842.7519798278809, + 817.1200156211853, + 820.1599717140198, + 829.472005367279, + 826.3999819755554, + 836.575984954834, + 846.8480110168457, + 869.376003742218, + 824.3200182914734, + 826.4960050582886, + 832.5440287590027, + 826.4639973640442, + 842.7839875221252, + 825.3759741783142, + 829.4399976730347, + 887.4559998512268, + 825.3120183944702, + 889.8559808731079, + 876.5439987182617, + 859.1679930686951, + 887.8719806671143, + 885.7920169830322, + 835.42400598526, + 852.0960211753845, + 857.0560216903687, + 834.5280289649963, + 837.6320004463196, + 832.5120210647583, + 806.0479760169983, + 848.1280207633972, + 870.4320192337036, + 852.9599905014038, + 833.5679769515991, + 843.7439799308777, + 905.8560132980347, + 852.9919981956482, + 871.4240193367004, + 836.6079926490784, + 823.3280181884766, + 873.4719753265381, + 888.8000249862671, + 866.5919899940491, + 789.5359992980957, + 834.496021270752, + 914.3999814987183, + 808.5759878158569, + 826.367974281311, + 898.0479836463928, + 827.3919820785522, + 843.5840010643005, + 849.951982498169, + 831.4239978790283, + 875.5199909210205, + 893.8559889793396, + 876.5439987182617, + 829.472005367279, + 861.2480163574219, + 881.6959857940674, + 886.7200016975403, + 826.3999819755554, + 840.6720161437988, + 857.1199774742126, + 869.3439960479736, + 836.6079926490784, + 829.4399976730347, + 856.1279773712158, + 884.768009185791, + 872.4480271339417, + 823.1359720230103, + 885.7600092887878, + 856.0960292816162, + 837.6320004463196, + 896.0959911346436, + 890.9119963645935, + 885.7600092887878, + 651.2960195541382, + 646.7519998550415, + 901.1200070381165, + 880.3520202636719, + 895.7759737968445, + 849.951982498169, + 820.2239871025085, + 837.6320004463196, + 834.496021270752, + 785.4080200195312, + 857.088029384613, + 879.7119855880737, + 868.511974811554, + 886.784017086029, + 797.7280020713806, + 837.6320004463196, + 888.8319730758667, + 850.9759902954102, + 911.3280177116394, + 868.3840036392212, + 829.5999765396118, + 891.9360041618347, + 866.27197265625, + 882.6879858970642, + 693.2799816131592, + 677.8879761695862, + 867.35999584198, + 874.4959831237793, + 885.7920169830322, + 888.8319730758667, + 854.9759984016418, + 880.6399703025818, + 880.5440068244934, + 889.9199962615967, + 841.7279720306396, + 893.9840197563171, + 834.5919847488403, + 860.0000143051147, + 828.2880187034607, + 824.3520259857178, + 835.5519771575928, + 878.5279989242554, + 893.3759927749634, + 828.4159898757935, + 833.5360288619995, + 877.5680065155029, + 874.4639754295349, + 832.5440287590027, + 864.3199801445007, + 840.7040238380432, + 835.5839848518372, + 873.4719753265381, + 677.8879761695862, + 869.376003742218, + 894.0160274505615, + 877.5680065155029, + 870.4000115394592, + 706.3040137290955, + 893.9840197563171, + 877.9839873313904, + 829.4399976730347, + 839.6480083465576, + 834.5919847488403, + 898.0479836463928, + 875.5519986152649, + 870.464026927948, + 817.2159790992737, + 826.3999819755554, + 871.4240193367004, + 867.35999584198, + 894.976019859314, + 837.664008140564, + 851.967990398407, + 851.1360287666321, + 837.664008140564, + 841.7279720306396, + 838.6560082435608, + 905.3120017051697, + 843.8079953193665, + 820.2559947967529, + 882.7199935913086, + 835.5839848518372, + 879.647970199585, + 835.6159925460815, + 674.8160123825073, + 836.7040157318115, + 833.5679769515991, + 831.4560055732727, + 882.5280070304871, + 824.3200182914734, + 878.495991230011, + 830.4960131645203, + 843.7439799308777, + 880.6399703025818, + 886.8799805641174, + 836.2879753112793, + 866.3039803504944, + 846.8480110168457, + 850.9439826011658, + 842.7519798278809, + 692.2879815101624, + 811.0399842262268, + 869.376003742218, + 876.5439987182617, + 827.2960186004639, + 716.543972492218, + 883.7440013885498, + 868.3519959449768, + 888.9279961585999, + 830.4319977760315, + 865.2799725532532, + 822.3040103912354, + 826.367974281311, + 833.6640000343323 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 188.7200027704239, + "candidate_gbps": 773.1275850895977, + "candidate_peak_mib": 42.103515625, + "provider_us": 184.25600230693817, + "provider_gbps": 791.8582742121392, + "provider_peak_mib": 168.1025390625, + "candidate_backward_us": 1395.1520323753357, + "candidate_backward_gbps": 104.57974228915255, + "candidate_backward_peak_mib": 102.67578125, + "provider_backward_us": 846.8480110168457, + "provider_backward_gbps": 172.29141251073642, + "provider_backward_peak_mib": 126.0615234375 + }, + { + "op": "final_adaln_out", + "case": "S=32768", + "backend": "H3FinalAdaLNOutCudaOp", + "bytes": 762445824, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 406.43200278282166, + 1319.2640542984009, + 1312.127947807312, + 2313.6959075927734, + 2309.4398975372314, + 2702.2719383239746, + 2698.4639167785645, + 449.0239918231964, + 403.3600091934204, + 1297.4400520324707, + 2279.520034790039, + 2282.8478813171387, + 2280.6079387664795, + 948.1279850006104, + 951.8399834632874, + 447.1679925918579, + 398.0480134487152, + 2044.3840026855469, + 2043.3919429779053, + 754.688024520874, + 2716.480016708374, + 2696.4480876922607, + 2698.6560821533203, + 2703.9999961853027, + 2703.455924987793, + 2076.927900314331, + 2714.7200107574463, + 2715.1360511779785, + 2078.04799079895, + 2713.2160663604736, + 2716.7999744415283, + 2063.807964324951, + 2718.2719707489014, + 1588.6080265045166, + 2717.8239822387695, + 1592.8959846496582, + 2720.031976699829, + 1590.8160209655762, + 2716.032028198242, + 1581.5999507904053, + 2718.2719707489014, + 1566.815972328186, + 2715.872049331665, + 1589.568018913269, + 2718.303918838501, + 2711.1361026763916, + 2714.8799896240234, + 2712.735891342163, + 2719.3920612335205, + 2713.9201164245605, + 2719.935894012451, + 2713.792085647583, + 2720.992088317871, + 2710.911989212036, + 2721.9200134277344, + 2714.047908782959, + 2718.0159091949463, + 2715.5840396881104, + 2752.4800300598145, + 446.368008852005, + 399.1360068321228, + 451.07200741767883, + 400.41598677635193, + 444.7680115699768, + 400.89601278305054, + 447.84000515937805, + 399.9359905719757, + 450.111985206604, + 397.63200283050537, + 449.3759870529175, + 405.40799498558044, + 448.2559859752655, + 396.8000113964081, + 452.60798931121826, + 413.4399890899658, + 450.20800828933716, + 405.0559997558594, + 451.4240026473999, + 410.68801283836365, + 446.368008852005, + 400.2240002155304, + 445.5359876155853, + 401.0240137577057, + 449.15199279785156, + 401.18399262428284, + 448.35200905799866, + 399.1039991378784, + 462.7839922904968, + 400.0000059604645, + 447.7440118789673, + 406.5600037574768, + 895.2959775924683, + 845.3440070152283, + 475.9039878845215, + 401.8239974975586, + 450.080007314682, + 401.2799859046936, + 449.5680034160614, + 400.9599983692169, + 446.04799151420593, + 399.9359905719757, + 448.2559859752655, + 397.7920114994049, + 447.90399074554443, + 399.55198764801025, + 446.3360011577606, + 397.95199036598206, + 447.32800126075745, + 410.0480079650879, + 451.35998725891113, + 412.9599928855896, + 429.8880100250244, + 425.6640076637268, + 451.231986284256, + 412.4799966812134, + 467.3599898815155, + 402.97600626945496, + 457.3119878768921, + 400.5120098590851, + 456.60799741744995, + 402.3039937019348, + 434.4319999217987, + 430.4639995098114, + 452.5119960308075, + 399.3600010871887, + 452.86399126052856, + 402.5920033454895, + 461.9840085506439, + 401.0240137577057, + 456.8000137805939, + 398.9439904689789, + 451.231986284256, + 400.5439877510071, + 456.7359983921051, + 401.5040099620819, + 451.1359930038452, + 399.29598569869995, + 458.40001106262207, + 399.80798959732056, + 450.52799582481384, + 398.6240029335022, + 457.8559994697571, + 399.80798959732056, + 455.6480050086975, + 407.9360067844391, + 454.17600870132446, + 399.00800585746765, + 449.91999864578247, + 401.15201473236084, + 408.7679982185364, + 744.2880272865295, + 672.3840236663818, + 568.7040090560913, + 568.8959956169128, + 676.7359972000122, + 575.1680135726929, + 674.0800142288208, + 674.1759777069092, + 674.5920181274414, + 676.3520240783691, + 672.9599833488464, + 676.1599779129028, + 676.6719818115234, + 638.1440162658691, + 635.8720064163208, + 569.8239803314209, + 982.8159809112549, + 640.7359838485718, + 988.1600141525269, + 641.4080262184143, + 983.7759733200073, + 450.20800828933716, + 397.69598841667175, + 975.0720262527466, + 971.0080027580261, + 2713.855981826782, + 1904.7679901123047, + 1000.3520250320435, + 2704.576015472412, + 2713.632106781006, + 1903.167963027954, + 2712.0959758758545, + 2716.0000801086426, + 2714.560031890869, + 2331.808090209961, + 2337.88800239563, + 2705.9199810028076, + 626.1759996414185, + 2703.903913497925, + 449.75998997688293, + 414.2400026321411, + 451.7120122909546, + 400.736004114151, + 471.16801142692566, + 418.08000206947327, + 481.85598850250244, + 422.9759871959686, + 448.7040042877197, + 411.6800129413605, + 452.2239863872528 + ], + "provider": [ + 801.7920255661011, + 762.2399926185608, + 764.9919986724854, + 764.4799947738647, + 762.0159983634949, + 764.8000121116638, + 767.0080065727234, + 876.5439987182617, + 800.4800081253052, + 762.5920176506042, + 764.5760178565979, + 767.0080065727234, + 760.5440020561218, + 766.6559815406799, + 764.9279832839966, + 872.3840117454529, + 813.9200210571289, + 764.3200159072876, + 762.9439830780029, + 762.4959945678711, + 766.5280103683472, + 772.8319764137268, + 770.4319953918457, + 764.959990978241, + 766.3040161132812, + 767.0400142669678, + 769.0560221672058, + 766.1120295524597, + 770.3040242195129, + 764.8320198059082, + 764.8959755897522, + 766.2400007247925, + 766.7199969291687, + 762.8800272941589, + 762.3680233955383, + 766.3040161132812, + 770.8479762077332, + 764.352023601532, + 763.0079984664917, + 768.127977848053, + 766.2720084190369, + 766.1759853363037, + 766.3679718971252, + 764.2880082130432, + 766.7520046234131, + 764.3839716911316, + 770.7840204238892, + 764.6399736404419, + 764.735996723175, + 762.175977230072, + 768.0960297584534, + 762.4959945678711, + 762.8160119056702, + 762.0800137519836, + 764.1599774360657, + 764.8959755897522, + 768.9279913902283, + 761.9519829750061, + 768.8320279121399, + 878.5600066184998, + 798.6559867858887, + 877.951979637146, + 798.6559867858887, + 873.2159733772278, + 804.5120239257812, + 877.9199719429016, + 798.8799810409546, + 872.3520040512085, + 805.9840202331543, + 871.8720078468323, + 806.8479895591736, + 877.3760199546814, + 802.1439909934998, + 876.2239813804626, + 782.5599908828735, + 875.0079870223999, + 816.4160251617432, + 872.223973274231, + 803.5839796066284, + 883.4879994392395, + 808.031976222992, + 878.2079815864563, + 814.9440288543701, + 872.7359771728516, + 802.1439909934998, + 872.7999925613403, + 800.1599907875061, + 874.4959831237793, + 818.5279965400696, + 874.9120235443115, + 815.6160116195679, + 1209.0879678726196, + 1115.4240369796753, + 874.9439716339111, + 799.5839715003967, + 874.4639754295349, + 803.4240007400513, + 760.0640058517456, + 803.2960295677185, + 871.7120289802551, + 797.0560193061829, + 875.3920197486877, + 801.0560274124146, + 873.9519715309143, + 799.0080118179321, + 875.8400082588196, + 799.9680042266846, + 873.3760118484497, + 799.5200157165527, + 882.5600147247314, + 798.8799810409546, + 871.936023235321, + 802.8159737586975, + 878.6240220069885, + 802.5919795036316, + 887.7120018005371, + 800.5120158195496, + 889.4400000572205, + 820.1280236244202, + 855.4559946060181, + 799.5839715003967, + 878.9119720458984, + 800.9600043296814, + 880.0960183143616, + 836.6720080375671, + 877.7920007705688, + 805.0879836082458, + 877.344012260437, + 819.4239735603333, + 883.3919763565063, + 803.2640218734741, + 880.4799914360046, + 803.0080199241638, + 880.3520202636719, + 804.9280047416687, + 882.5920224189758, + 803.9680123329163, + 878.5600066184998, + 802.5280237197876, + 885.8240246772766, + 801.3439774513245, + 880.511999130249, + 805.2800297737122, + 877.8240084648132, + 800.3519773483276, + 883.9359879493713, + 804.7999739646912, + 882.9759955406189, + 803.5839796066284, + 878.7199854850769, + 769.0879702568054, + 760.1600289344788, + 766.3360238075256, + 766.2079930305481, + 768.3839797973633, + 769.0879702568054, + 766.3999795913696, + 764.2880082130432, + 764.7039890289307, + 760.9279751777649, + 769.0879702568054, + 768.8320279121399, + 768.671989440918, + 758.6560249328613, + 764.6719813346863, + 765.0560140609741, + 760.0640058517456, + 764.0640139579773, + 762.4639868736267, + 770.3679800033569, + 773.1840014457703, + 877.951979637146, + 803.2320141792297, + 762.3999714851379, + 765.0240063667297, + 764.1280293464661, + 768.8320279121399, + 764.959990978241, + 771.1039781570435, + 764.5120024681091, + 766.975998878479, + 766.2400007247925, + 766.6559815406799, + 764.9279832839966, + 770.7840204238892, + 766.7520046234131, + 768.6079740524292, + 766.8160200119019, + 766.8480277061462, + 886.6559863090515, + 801.6960024833679, + 880.4799914360046, + 804.3199777603149, + 890.1439905166626, + 803.6159873008728, + 891.7760252952576, + 797.5040078163147, + 878.4319758415222, + 817.6320195198059, + 881.3760280609131 + ], + "candidate_backward": [ + 2507.8399181365967, + 5161.1199378967285, + 7166.048049926758, + 6152.351856231689, + 7003.232002258301, + 5903.42378616333, + 7601.344108581543, + 2509.82403755188, + 2554.8479557037354, + 5101.7279624938965, + 7073.056221008301, + 6095.935821533203, + 5826.623916625977, + 7546.080112457275, + 5827.712059020996, + 2578.399896621704, + 2484.2560291290283, + 7344.2559242248535, + 6716.479778289795, + 6711.264133453369, + 6713.535785675049, + 6726.719856262207, + 8816.767692565918, + 8331.328392028809, + 2475.1360416412354, + 9347.264289855957, + 8729.696273803711, + 8704.095840454102, + 8705.151557922363, + 9349.21646118164, + 9350.303649902344, + 9355.327606201172, + 8743.00765991211, + 8716.480255126953, + 8723.520278930664, + 8734.848022460938, + 8750.176429748535, + 8727.744102478027, + 8730.78441619873, + 8720.576286315918, + 8734.75170135498, + 8709.280014038086, + 8716.383934020996, + 8744.095802307129, + 6633.600234985352, + 9342.240333557129, + 9352.352142333984, + 9349.21646118164, + 9353.280067443848, + 9343.135833740234, + 9350.208282470703, + 9359.456062316895, + 9362.848281860352, + 9366.592407226562, + 9350.144386291504, + 9352.288246154785, + 9356.351852416992, + 9350.144386291504, + 9356.32038116455, + 2516.9920921325684, + 2506.7200660705566, + 2514.8160457611084, + 2530.303955078125, + 2536.479949951172, + 2505.631923675537, + 2530.2720069885254, + 2474.976062774658, + 2572.2880363464355, + 2498.528003692627, + 2484.2240810394287, + 2543.6160564422607, + 2565.0880336761475, + 2500.7998943328857, + 2515.80810546875, + 2573.280096054077, + 2542.5920486450195, + 2501.6000270843506, + 2589.6639823913574, + 2504.5759677886963, + 2539.520025253296, + 2585.599899291992, + 2472.991943359375, + 2574.336051940918, + 2554.879903793335, + 2512.3839378356934, + 2505.6960582733154, + 2512.864112854004, + 2595.8080291748047, + 2476.032018661499, + 2499.583959579468, + 2511.8720531463623, + 2841.599941253662, + 2750.52809715271, + 2707.2958946228027, + 2481.8880558013916, + 2618.367910385132, + 2487.391948699951, + 2685.9519481658936, + 2465.791940689087, + 2471.679925918579, + 2497.5359439849854, + 2468.832015991211, + 2559.999942779541, + 2541.5680408477783, + 2495.4559803009033, + 2525.1200199127197, + 2513.9200687408447, + 2517.983913421631, + 2621.4399337768555, + 2645.983934402466, + 2618.367910385132, + 2528.223991394043, + 2525.2161026000977, + 2563.0080699920654, + 2491.8720722198486, + 2519.455909729004, + 2488.1598949432373, + 2567.13604927063, + 2554.7521114349365, + 2553.8558959960938, + 2506.7520141601562, + 2582.047939300537, + 2504.2240619659424, + 2569.216012954712, + 2508.8319778442383, + 2578.432083129883, + 2557.9519271850586, + 2591.7439460754395, + 2511.7759704589844, + 2576.4479637145996, + 2492.608070373535, + 2517.983913421631, + 2516.9920921325684, + 2567.13604927063, + 2461.695909500122, + 2528.287887573242, + 2541.5360927581787, + 2586.5280628204346, + 2535.423994064331, + 2457.6001167297363, + 2573.3120441436768, + 2521.0559368133545, + 2502.2079944610596, + 2533.344030380249, + 2527.2319316864014, + 2631.711959838867, + 2566.1120414733887, + 2531.2960147857666, + 2541.599988937378, + 2512.8960609436035, + 7023.615837097168, + 3231.7440509796143, + 4989.984035491943, + 4979.743957519531, + 4983.776092529297, + 6944.863796234131, + 6942.65604019165, + 3238.9121055603027, + 4948.160171508789, + 7096.255779266357, + 4955.135822296143, + 5113.823890686035, + 6883.264064788818, + 3119.0719604492188, + 3058.6559772491455, + 5123.263835906982, + 5132.351875305176, + 5558.271884918213, + 5615.7121658325195, + 5619.840145111084, + 5545.951843261719, + 2462.7199172973633, + 2485.215902328491, + 4148.416042327881, + 5489.664077758789, + 6836.319923400879, + 8573.023796081543, + 9340.319633483887, + 5489.7918701171875, + 6831.295967102051, + 8527.071952819824, + 9339.872360229492, + 9342.111587524414, + 8938.655853271484, + 8978.560447692871, + 8984.640121459961, + 9339.967727661133, + 9344.063758850098, + 2473.0560779571533, + 2580.0960063934326, + 2568.3839321136475, + 2578.336000442505, + 2500.7359981536865, + 2588.6080265045166, + 2518.143892288208, + 2551.2959957122803, + 2511.807918548584, + 2553.823947906494, + 2555.8719635009766, + 2476.032018661499 + ], + "provider_backward": [ + 4461.27986907959, + 5841.663837432861, + 7606.400012969971, + 7607.52010345459, + 7670.8478927612305, + 7672.99222946167, + 7673.952102661133, + 7670.91178894043, + 6567.039966583252, + 6561.855792999268, + 7531.519889831543, + 7545.82405090332, + 7622.7521896362305, + 7627.903938293457, + 5170.2399253845215, + 5130.271911621094, + 4455.2321434021, + 6577.919960021973, + 6781.055927276611, + 6756.383895874023, + 8722.5923538208, + 8738.847732543945, + 8876.22356414795, + 8886.336326599121, + 7348.991870880127, + 9392.127990722656, + 9399.200439453125, + 9413.6962890625, + 9394.111633300781, + 8804.320335388184, + 8760.416030883789, + 9404.607772827148, + 8316.831588745117, + 9401.344299316406, + 8286.208152770996, + 9405.440330505371, + 8302.656173706055, + 9404.255867004395, + 8303.520202636719, + 9397.279739379883, + 8299.519538879395, + 9405.50422668457, + 8310.815811157227, + 9413.6962890625, + 9408.543586730957, + 9410.623550415039, + 9389.984130859375, + 9398.303985595703, + 9396.35181427002, + 9400.287628173828, + 9404.41608428955, + 9406.559944152832, + 9400.416374206543, + 9394.207954406738, + 9403.424263000488, + 9399.328231811523, + 9389.120101928711, + 9388.128280639648, + 9387.07160949707, + 4467.807769775391, + 4456.1920166015625, + 4459.968090057373, + 4470.911979675293, + 4468.863964080811, + 4475.039958953857, + 4469.823837280273, + 4466.879844665527, + 4456.511974334717, + 4464.416027069092, + 4463.712215423584, + 4456.575870513916, + 4468.832015991211, + 4458.176136016846, + 4468.863964080811, + 4785.215854644775, + 5010.464191436768, + 4458.208084106445, + 4457.632064819336, + 4457.248210906982, + 4475.103855133057, + 4464.7040367126465, + 4469.888210296631, + 4464.640140533447, + 4475.039958953857, + 4465.792179107666, + 4460.256099700928, + 4456.511974334717, + 4472.671985626221, + 4454.368114471436, + 4469.98405456543, + 4455.615997314453, + 4467.872142791748, + 4501.503944396973, + 4459.4879150390625, + 4466.815948486328, + 4474.815845489502, + 4460.288047790527, + 4462.7838134765625, + 4456.480026245117, + 4464.7040367126465, + 4465.888023376465, + 4466.623783111572, + 4465.375900268555, + 4458.208084106445, + 4468.512058258057, + 4458.240032196045, + 4462.33606338501, + 5305.408000946045, + 4475.679874420166, + 4479.040145874023, + 4463.967800140381, + 4461.696147918701, + 4501.664161682129, + 4462.368011474609, + 4460.7038497924805, + 4479.2962074279785, + 4458.432197570801, + 4487.232208251953, + 4468.448162078857, + 4466.879844665527, + 4464.7040367126465, + 4454.239845275879, + 4456.543922424316, + 4457.215785980225, + 4463.359832763672, + 4462.751865386963, + 4466.495990753174, + 4478.784084320068, + 4495.0079917907715, + 4469.535827636719, + 4460.671901702881, + 4462.368011474609, + 4470.49617767334, + 4458.367824554443, + 4460.383892059326, + 4458.335876464844, + 4456.480026245117, + 4468.89591217041, + 4453.63187789917, + 4468.512058258057, + 4465.792179107666, + 4454.495906829834, + 4465.824127197266, + 4468.544006347656, + 4464.767932891846, + 4458.5280418396, + 4471.871852874756, + 4454.239845275879, + 4464.416027069092, + 4472.671985626221, + 5325.632095336914, + 5326.015949249268, + 7161.888122558594, + 7164.000034332275, + 7173.151969909668, + 9247.83992767334, + 4457.536220550537, + 4776.735782623291, + 7151.7438888549805, + 7148.543834686279, + 7147.456169128418, + 9083.968162536621, + 9164.863586425781, + 4464.640140533447, + 5603.392124176025, + 5600.224018096924, + 7329.855918884277, + 7334.815979003906, + 7677.984237670898, + 7761.8560791015625, + 7670.720100402832, + 7750.592231750488, + 4462.30411529541, + 4880.127906799316, + 9386.079788208008, + 7679.967880249023, + 9402.400016784668, + 8644.576072692871, + 9389.056205749512, + 7684.031963348389, + 9387.136459350586, + 8645.600318908691, + 9378.87954711914, + 9063.520431518555, + 9393.152236938477, + 9386.943817138672, + 9013.18359375, + 9405.471801757812, + 4455.167770385742, + 4481.728076934814, + 7643.2318687438965, + 4459.680080413818, + 4459.712028503418, + 4466.432094573975, + 4459.296226501465, + 4483.8080406188965, + 4455.615997314453, + 4457.503795623779, + 4476.5119552612305, + 4467.872142791748 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 458.1280052661896, + "candidate_gbps": 1664.2637324845277, + "candidate_peak_mib": 336.2021484375, + "provider_us": 799.2640137672424, + "provider_gbps": 953.9348836766666, + "provider_peak_mib": 1344.0615234375, + "candidate_backward_us": 2590.7039642333984, + "candidate_backward_gbps": 294.30063586042, + "candidate_backward_peak_mib": 397.10400390625, + "provider_backward_us": 4485.520124435425, + "provider_backward_gbps": 169.97935642880796, + "provider_backward_peak_mib": 1008.03076171875 + }, + { + "op": "final_adaln_out", + "case": "S=131072", + "backend": "H3FinalAdaLNOutCudaOp", + "bytes": 2876375040, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "timing_samples_us": { + "candidate": [ + 1413.1519794464111, + 3560.4801177978516, + 3586.8160724639893, + 2163.167953491211, + 2161.6320610046387, + 2898.8800048828125, + 2892.2879695892334, + 1419.0399646759033, + 1598.8479852676392, + 3568.608045578003, + 3567.5199031829834, + 3563.488006591797, + 3567.1041011810303, + 3566.335916519165, + 1425.6960153579712, + 3566.080093383789, + 1421.6320514678955, + 3561.055898666382, + 1426.5919923782349, + 3565.023899078369, + 3560.2879524230957, + 3559.904098510742, + 3557.1200847625732, + 3556.4799308776855, + 1945.6000328063965, + 1923.2319593429565, + 3566.272020339966, + 1315.071940422058, + 1244.5119619369507, + 1322.9119777679443, + 1248.1919527053833, + 1332.3520421981812, + 1244.7999715805054, + 1311.5839958190918, + 1269.152045249939, + 1308.2560300827026, + 1245.9839582443237, + 1521.2160348892212, + 1514.5920515060425, + 1490.9440279006958, + 1483.8080406188965, + 1424.5760440826416, + 1764.415979385376, + 1698.0479955673218, + 1687.008023262024, + 2663.5520458221436, + 1483.5840463638306, + 1423.5199689865112, + 1677.3439645767212, + 1425.2480268478394, + 1482.0159673690796, + 1308.0639839172363, + 1245.8560466766357, + 1420.7040071487427, + 2898.655891418457, + 1466.0799503326416, + 3251.039981842041, + 1328.6399841308594, + 1247.9360103607178, + 1484.5759868621826, + 1489.8879528045654, + 1305.408000946045, + 1254.7199726104736, + 1307.4239492416382, + 1243.008017539978, + 1427.2960424423218, + 2188.4799003601074, + 1456.1920166015625, + 2352.384090423584, + 1304.3839931488037, + 1248.1600046157837, + 2167.2959327697754, + 2166.7520999908447, + 1427.8719425201416, + 2313.663959503174, + 2194.6239471435547, + 2191.0719871520996, + 2161.2799167633057, + 2160.6080532073975, + 3565.5040740966797, + 1887.4560594558716, + 1896.83198928833, + 2689.5039081573486, + 3563.7118816375732, + 3551.2959957122803, + 1355.7440042495728, + 1266.6560411453247, + 1341.4720296859741, + 1262.4640464782715, + 1320.7039833068848, + 1260.5119943618774, + 1305.0240278244019, + 1250.3679990768433, + 1421.247959136963, + 1266.752004623413, + 1317.471981048584, + 1281.8880081176758, + 1312.8639459609985, + 1268.7040567398071, + 1313.3440017700195, + 1281.6319465637207, + 1339.1040563583374, + 1279.1039943695068, + 1373.568058013916, + 1278.3679962158203, + 1319.0720081329346, + 1268.5760259628296, + 1317.3760175704956, + 1264.0000581741333, + 1308.5440397262573, + 1260.032057762146, + 1315.5839443206787, + 1271.3279724121094, + 1321.984052658081, + 1279.3920040130615, + 1313.9519691467285, + 1279.9999713897705, + 1329.3440341949463, + 1266.271948814392, + 1321.2480545043945, + 1259.8400115966797, + 1324.7040510177612, + 1262.9120349884033, + 1325.5679607391357, + 1262.943983078003, + 1334.0480327606201, + 1261.1520290374756, + 1338.3359909057617, + 1264.9279832839966, + 1315.8719539642334, + 1259.5839500427246, + 1325.2480030059814, + 1259.775996208191, + 1314.687967300415, + 1269.760012626648, + 1314.303994178772, + 1265.6960487365723, + 1312.9279613494873, + 1267.0079469680786, + 1312.4480247497559, + 1259.2320442199707, + 1310.655951499939, + 1258.7840557098389, + 1312.991976737976, + 1257.599949836731, + 1305.8240413665771, + 1259.4239711761475, + 1310.3359937667847, + 1256.00004196167, + 1335.8720541000366, + 1280.7040214538574, + 1321.3119506835938, + 1259.775996208191, + 1327.2960186004639, + 1260.9280347824097, + 1333.8559865951538, + 1264.9600505828857, + 1340.1600122451782, + 1263.7120485305786, + 1340.6399488449097, + 1261.9520425796509, + 1326.97594165802, + 1260.6079578399658, + 1328.8320302963257, + 1260.8959674835205, + 1319.551944732666, + 1259.071946144104, + 1325.2480030059814, + 1262.719988822937, + 1317.9839849472046, + 1261.9839906692505, + 1318.6240196228027, + 1266.9440507888794, + 1314.3359422683716, + 1263.424038887024, + 1320.032000541687, + 1254.1760206222534, + 1313.215970993042, + 1266.495943069458, + 1308.3200454711914, + 1261.6959810256958, + 1329.6639919281006, + 1289.5359992980957, + 1315.1999711990356, + 1270.4960107803345, + 1328.544020652771, + 1263.0079984664917, + 1329.632043838501, + 1264.4799947738647, + 1321.4399814605713, + 1269.5679664611816, + 1338.879942893982, + 1262.112021446228, + 1329.0239572525024, + 1259.775996208191, + 1326.0480165481567, + 1264.2560005187988, + 1338.271975517273, + 1259.4879865646362, + 1323.9680528640747 + ], + "provider": [ + 4002.079963684082, + 5354.30383682251, + 2903.968095779419, + 3936.8319511413574, + 3978.271961212158, + 4716.479778289795, + 4719.488143920898, + 3434.720039367676, + 5384.736061096191, + 5383.743762969971, + 5392.864227294922, + 5396.639823913574, + 5396.063804626465, + 5371.007919311523, + 5369.855880737305, + 5382.3041915893555, + 5372.191905975342, + 5393.248081207275, + 5378.464221954346, + 5392.735958099365, + 5350.944042205811, + 5363.552093505859, + 3754.688024520874, + 5391.9677734375, + 5392.576217651367, + 5388.832092285156, + 5395.5841064453125, + 3019.7439193725586, + 2934.2079162597656, + 3021.3119983673096, + 2937.5998973846436, + 3027.1360874176025, + 3264.1921043395996, + 3016.2880420684814, + 2941.4079189300537, + 3012.864112854004, + 2931.2639236450195, + 4052.95991897583, + 3298.975944519043, + 3971.776008605957, + 4047.679901123047, + 4208.064079284668, + 3581.631898880005, + 3519.1359519958496, + 3495.6159591674805, + 4485.375881195068, + 4483.871936798096, + 4143.519878387451, + 3499.552011489868, + 4161.600112915039, + 4162.303924560547, + 4035.3918075561523, + 2931.839942932129, + 4760.32018661499, + 5390.880107879639, + 5355.648040771484, + 2932.7359199523926, + 3011.807918548584, + 2937.8559589385986, + 3016.479969024658, + 2930.720090866089, + 2994.1439628601074, + 2936.5758895874023, + 3007.967948913574, + 2932.800054550171, + 3880.5758953094482, + 4014.7199630737305, + 4155.776023864746, + 4172.544002532959, + 3015.8400535583496, + 4150.815963745117, + 4158.847808837891, + 4969.183921813965, + 4054.0480613708496, + 4131.584167480469, + 4015.359878540039, + 4001.5358924865723, + 4826.848030090332, + 4826.272010803223, + 5360.991954803467, + 5347.743988037109, + 5361.536026000977, + 5343.808174133301, + 5373.311996459961, + 5377.984046936035, + 3038.304090499878, + 2965.503931045532, + 3029.184103012085, + 2944.000005722046, + 3047.1999645233154, + 2946.2718963623047, + 3002.239942550659, + 2942.239999771118, + 3058.624029159546, + 2952.7039527893066, + 3017.983913421631, + 2963.104009628296, + 3015.3279304504395, + 2950.7200717926025, + 3010.5600357055664, + 2949.0880966186523, + 3029.0560722351074, + 2961.7600440979004, + 3015.5200958251953, + 2954.655885696411, + 2998.9120960235596, + 2954.3681144714355, + 3008.1279277801514, + 2961.632013320923, + 2986.36794090271, + 2965.503931045532, + 3025.631904602051, + 2974.816083908081, + 3010.688066482544, + 2946.2718963623047, + 3009.5999240875244, + 2944.4799423217773, + 3022.752046585083, + 2950.4640102386475, + 3017.024040222168, + 2949.2480754852295, + 3009.183883666992, + 2946.9120502471924, + 3013.184070587158, + 2948.512077331543, + 3021.7599868774414, + 2951.4880180358887, + 3033.440113067627, + 2954.943895339966, + 3017.6000595092773, + 2946.943998336792, + 3021.6000080108643, + 2946.399927139282, + 3013.5679244995117, + 2959.264039993286, + 3005.311965942383, + 2964.224100112915, + 3009.984016418457, + 2958.336114883423, + 3013.823986053467, + 2961.3120555877686, + 3014.8799419403076, + 2956.511974334717, + 3012.928009033203, + 2945.5039501190186, + 2987.6160621643066, + 2951.263904571533, + 3011.712074279785, + 2953.3441066741943, + 3073.1520652770996, + 2953.9198875427246, + 2985.408067703247, + 2949.023962020874, + 3017.6000595092773, + 2947.4239349365234, + 3012.768030166626, + 2944.6399211883545, + 3024.9600410461426, + 2944.1280364990234, + 3036.0960960388184, + 2944.4479942321777, + 3028.9599895477295, + 2946.6240406036377, + 3000.960111618042, + 2944.159984588623, + 3029.247999191284, + 2950.495958328247, + 3024.735927581787, + 2946.6559886932373, + 3030.143976211548, + 2949.824094772339, + 3012.1281147003174, + 2964.9600982666016, + 3008.6400508880615, + 2965.6639099121094, + 3011.45601272583, + 2952.5439739227295, + 2987.839937210083, + 2964.384078979492, + 2982.367992401123, + 2948.767900466919, + 3033.6320400238037, + 2952.7039527893066, + 3018.5599327087402, + 2950.495958328247, + 3018.3041095733643, + 2948.352098464966, + 3014.4639015197754, + 2946.0160732269287, + 3019.4239616394043, + 2957.7279090881348, + 3027.2960662841797, + 2942.68798828125, + 3002.880096435547, + 2947.200059890747, + 3017.4720287323, + 2946.784019470215, + 3024.6400833129883, + 2948.767900466919, + 3030.240058898926 + ], + "candidate_backward": [ + 15111.295700073242, + 16137.311935424805, + 8399.840354919434, + 14014.687538146973, + 17323.16780090332, + 7267.327785491943, + 16684.192657470703, + 13900.927543640137, + 19787.967681884766, + 7027.711868286133, + 20700.288772583008, + 20650.239944458008, + 20697.216033935547, + 20144.06394958496, + 20182.207107543945, + 20145.151138305664, + 20156.60858154297, + 20156.35108947754, + 20195.552825927734, + 20132.863998413086, + 21330.976486206055, + 19653.823852539062, + 21338.239669799805, + 19716.1922454834, + 21309.600830078125, + 21336.16065979004, + 17148.06365966797, + 7034.8801612854, + 7308.288097381592, + 7037.8241539001465, + 7043.680191040039, + 7032.832145690918, + 7111.743927001953, + 7057.375907897949, + 7065.567970275879, + 7030.784130096436, + 7024.640083312988, + 10861.568450927734, + 8941.66374206543, + 9983.00838470459, + 11869.215965270996, + 12076.255798339844, + 12096.735954284668, + 8686.623573303223, + 11189.311981201172, + 14594.112396240234, + 8743.935585021973, + 11706.368446350098, + 14083.168029785156, + 9781.15177154541, + 11372.447967529297, + 12407.679557800293, + 8910.880088806152, + 16476.255416870117, + 16462.047576904297, + 18568.607330322266, + 7037.951946258545, + 7037.951946258545, + 7063.5199546813965, + 7043.07222366333, + 7052.288055419922, + 7037.919998168945, + 7043.039798736572, + 7033.85591506958, + 14046.272277832031, + 11841.504096984863, + 11890.751838684082, + 13911.231994628906, + 12791.999816894531, + 16556.127548217773, + 13794.68822479248, + 12770.496368408203, + 18066.591262817383, + 12584.927558898926, + 12634.30404663086, + 12331.999778747559, + 12347.519874572754, + 16349.248886108398, + 7385.983943939209, + 17874.975204467773, + 16584.768295288086, + 18739.200592041016, + 18751.455307006836, + 20513.952255249023, + 20514.976501464844, + 16008.224487304688, + 7053.311824798584, + 7035.9039306640625, + 7035.935878753662, + 7026.656150817871, + 7033.8239669799805, + 7037.9838943481445, + 7031.839847564697, + 7033.85591506958, + 7039.968013763428, + 7033.8239669799805, + 7059.455871582031, + 7044.223785400391, + 7040.895938873291, + 7045.055866241455, + 7096.320152282715, + 7159.776210784912, + 7094.272136688232, + 7092.031955718994, + 7034.912109375, + 7041.024208068848, + 7038.976192474365, + 7031.807899475098, + 7036.928176879883, + 7030.784130096436, + 7029.727935791016, + 7035.9039306640625, + 7035.871982574463, + 7028.543949127197, + 7061.503887176514, + 7032.832145690918, + 7041.056156158447, + 7116.767883300781, + 7042.04797744751, + 7034.8801612854, + 7051.296234130859, + 7039.999961853027, + 7045.087814331055, + 7036.928176879883, + 7036.960124969482, + 7031.839847564697, + 7034.719944000244, + 7032.864093780518, + 7037.9838943481445, + 7039.936065673828, + 7037.919998168945, + 7039.008140563965, + 7054.336071014404, + 7119.840145111084, + 7044.064044952393, + 7031.839847564697, + 7036.928176879883, + 7034.848213195801, + 7046.080112457275, + 7039.999961853027, + 7037.9838943481445, + 7032.800197601318, + 7041.024208068848, + 7040.99178314209, + 7032.864093780518, + 7039.999961853027, + 7035.871982574463, + 7036.928176879883, + 7034.912109375, + 7133.18395614624, + 7128.064155578613, + 7037.951946258545, + 7045.152187347412, + 7062.528133392334, + 7042.01602935791, + 7033.88786315918, + 7043.07222366333, + 7036.896228790283, + 7030.752182006836, + 7041.984081268311, + 7039.103984832764, + 7034.016132354736, + 7038.976192474365, + 7035.9039306640625, + 7032.864093780518, + 7132.12776184082, + 7051.231861114502, + 7025.631904602051, + 7045.184135437012, + 7128.032207489014, + 7042.04797744751, + 7046.144008636475, + 7039.008140563965, + 7033.8239669799805, + 7032.896041870117, + 7035.871982574463, + 7035.871982574463, + 7036.896228790283, + 7043.039798736572, + 7041.024208068848, + 7032.832145690918, + 7118.815898895264, + 7081.952095031738, + 7048.3198165893555, + 7049.376010894775, + 7043.07222366333, + 7043.039798736572, + 7036.928176879883, + 7042.04797744751, + 7043.136119842529, + 7036.928176879883, + 7049.215793609619, + 7037.9838943481445, + 7037.951946258545, + 7034.8801612854, + 7037.9838943481445, + 7036.831855773926, + 7106.560230255127, + 7033.85591506958, + 7034.848213195801 + ], + "provider_backward": [ + 30227.519989013672, + 32814.97573852539, + 26003.456115722656, + 30913.503646850586, + 31009.695053100586, + 32618.49594116211, + 31867.29621887207, + 31859.872817993164, + 35336.22360229492, + 37207.87048339844, + 38951.969146728516, + 38550.97579956055, + 38927.36053466797, + 39141.47186279297, + 38639.6484375, + 38686.75231933594, + 38664.222717285156, + 39019.649505615234, + 39029.727935791016, + 38646.751403808594, + 38154.20913696289, + 38187.137603759766, + 38562.81661987305, + 38226.91345214844, + 39721.920013427734, + 39726.14288330078, + 17696.895599365234, + 17682.592391967773, + 21107.776641845703, + 17657.95135498047, + 17710.23941040039, + 17989.66407775879, + 17711.200714111328, + 18354.30335998535, + 17692.800521850586, + 17669.279098510742, + 17684.703826904297, + 17676.416397094727, + 23558.208465576172, + 23519.296646118164, + 23303.327560424805, + 23771.167755126953, + 30090.303421020508, + 29187.135696411133, + 25615.455627441406, + 25665.632247924805, + 23029.855728149414, + 25796.57554626465, + 29608.928680419922, + 20827.2647857666, + 24571.96807861328, + 28303.61557006836, + 23806.943893432617, + 32516.128540039062, + 34417.69790649414, + 38363.07144165039, + 20503.616333007812, + 17691.71142578125, + 17672.38426208496, + 17996.864318847656, + 17697.887420654297, + 17682.4951171875, + 17662.23907470703, + 17678.43246459961, + 27246.75178527832, + 25909.151077270508, + 33149.95193481445, + 23530.527114868164, + 28517.375946044922, + 35792.96112060547, + 28314.687728881836, + 28239.0079498291, + 24538.143157958984, + 28705.82389831543, + 32263.137817382812, + 31694.87953186035, + 27950.143814086914, + 29192.224502563477, + 34704.41436767578, + 34674.91149902344, + 34017.311096191406, + 36367.488861083984, + 34748.54278564453, + 34495.521545410156, + 36578.43017578125, + 38975.521087646484, + 17662.14370727539, + 17678.688049316406, + 17690.719604492188, + 17670.24040222168, + 17692.73567199707, + 17672.256469726562, + 17697.05581665039, + 17686.59210205078, + 17682.592391967773, + 17686.62452697754, + 17682.655334472656, + 17684.768676757812, + 17670.24040222168, + 17680.479049682617, + 17682.559967041016, + 17672.38426208496, + 17676.35154724121, + 17690.656661987305, + 17873.98338317871, + 17662.14370727539, + 17703.136444091797, + 17676.38397216797, + 17686.71989440918, + 17678.62319946289, + 17666.080474853516, + 17832.063674926758, + 17674.463272094727, + 17684.67140197754, + 17682.559967041016, + 17656.0001373291, + 17676.35154724121, + 17682.62481689453, + 17666.175842285156, + 17690.81687927246, + 17692.768096923828, + 17674.463272094727, + 17674.272537231445, + 17694.81658935547, + 17672.351837158203, + 17676.38397216797, + 17670.272827148438, + 17682.62481689453, + 17713.375091552734, + 17684.54360961914, + 17674.367904663086, + 17696.928024291992, + 17680.383682250977, + 17680.416107177734, + 17672.319412231445, + 17686.71989440918, + 17676.38397216797, + 17672.28889465332, + 17703.07159423828, + 17662.015914916992, + 17678.367614746094, + 17676.511764526367, + 17668.319702148438, + 17690.719604492188, + 17682.592391967773, + 17676.416397094727, + 17702.97622680664, + 17684.5760345459, + 17680.60874938965, + 17682.559967041016, + 17686.687469482422, + 17674.30305480957, + 17682.527542114258, + 17682.559967041016, + 17664.19219970703, + 17707.07130432129, + 17696.800231933594, + 17672.38426208496, + 17672.319412231445, + 17705.184936523438, + 17672.351837158203, + 17658.079147338867, + 17670.400619506836, + 17674.4327545166, + 17670.24040222168, + 17684.608459472656, + 17672.38426208496, + 17653.823852539062, + 17696.83265686035, + 17684.5760345459, + 17657.888412475586, + 17666.208267211914, + 17674.335479736328, + 17680.511474609375, + 17676.416397094727, + 17672.38426208496, + 17668.224334716797, + 17686.81526184082, + 17666.048049926758, + 17690.752029418945, + 17666.080474853516, + 17688.703536987305, + 17698.87924194336, + 17682.655334472656, + 17690.784454345703, + 17666.336059570312, + 17662.04833984375, + 17668.19190979004, + 17694.78416442871, + 17666.175842285156, + 17676.47933959961, + 17682.559967041016, + 17684.703826904297, + 17678.560256958008, + 17674.400329589844, + 17664.064407348633, + 17668.256759643555, + 17670.400619506836, + 17666.112899780273, + 17688.735961914062 + ] + }, + "timing_order": [ + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ], + [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + ], + "candidate_us": 1320.367991924286, + "candidate_gbps": 2178.4646837795663, + "candidate_peak_mib": 1344.5771484375, + "provider_us": 3015.4240131378174, + "provider_gbps": 953.8874226205008, + "provider_peak_mib": 5376.0615234375, + "candidate_backward_us": 7045.120000839233, + "candidate_backward_gbps": 408.27906971880657, + "candidate_backward_peak_mib": 1406.60986328125, + "provider_backward_us": 17689.696311950684, + "provider_backward_gbps": 162.60171962685413, + "provider_backward_peak_mib": 4032.03076171875 + } + ] +} diff --git a/reports/experiments/h3-prior-art-b200/adaln_projection.json b/reports/experiments/h3-prior-art-b200/adaln_projection.json new file mode 100644 index 000000000..fb3f83d0e --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/adaln_projection.json @@ -0,0 +1,1347 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "adaln_projection", + "rl_kernel_commit": "c88691c6a100c04b03cb961bd7c0ee9df113e854", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc", + "mode_environment": { + "plain": {}, + "vllm": { + "CUBLAS_WORKSPACE_CONFIG": ":16:8", + "CUBLASLT_WORKSPACE_SIZE": "1" + }, + "sglang": {}, + "sglang_ieee": { + "TRITON_F32_DEFAULT": "ieee" + }, + "megatron_te_native": { + "CUBLASLT_WORKSPACE_SIZE": "0" + }, + "megatron_triton": {}, + "megatron_triton_ieee": { + "TRITON_F32_DEFAULT": "ieee" + } + } + }, + "quick": false, + "results": { + "plain": { + "diffusers[plain]": { + "source": "diffusers AdaLN projection, BF16 F.linear [plain]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 192, + "rowgrad_differ": 192 + }, + "257": { + "rows": 771, + "fwd_differ": 771, + "rowgrad_differ": 771 + }, + "2048": { + "rows": 6144, + "fwd_differ": 6144, + "rowgrad_differ": 6144 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 1024, + "rowgrad_differ": 2056 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 2232, + "first_failures": [ + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 771 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 771 + } + ] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 111.39199882745743, + "fwd_bwd_us": 768.0000066757202 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 100.11200234293938, + "fwd_bwd_us": 716.3679897785187 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 99.50399771332741, + "fwd_bwd_us": 725.7280051708221 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 0.0028456426244756147, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 110.73600128293037, + "fwd_bwd_us": 711.3760113716125 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.0022159011369222194, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 126.91199779510498, + "fwd_bwd_us": 728.3839881420135 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 650.7999897003174, + "fwd_bwd_us": 2345.616102218628 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3AdaLNProjectionCudaOp", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 1.1753830485451992e-06, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 100.60799866914749, + "fwd_bwd_us": 1937.3120069503784 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 1.4082949304861482e-06, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 100.63999891281128, + "fwd_bwd_us": 2276.2240171432495 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 1.4082949304861482e-06, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 100.47999769449234, + "fwd_bwd_us": 2114.735960960388 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 1.5823117650412027e-06, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 624.7680187225342, + "fwd_bwd_us": 19274.608612060547 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 1.374057475644507e-06, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 2569.6959495544434, + "fwd_bwd_us": 77446.19369506836 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 1.6500752601653308e-06, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 21261.792182922363, + "fwd_bwd_us": 626991.6076660156 + } + } + } + }, + "vllm": { + "diffusers[vllm]": { + "source": "diffusers AdaLN projection, BF16 F.linear [vllm]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 100.12800246477127, + "fwd_bwd_us": 825.0080049037933 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 99.85600039362907, + "fwd_bwd_us": 681.0559928417206 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 100.25600343942642, + "fwd_bwd_us": 677.5520145893097 + }, + "64": { + "fwd_rel_err": 0.003117907166490413, + "grad_rel_err": { + "temb": 0.002934168727885534, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 110.96000298857689, + "fwd_bwd_us": 679.85600233078 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.002284836389545336, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 126.51199847459793, + "fwd_bwd_us": 722.4960029125214 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 635.77601313591, + "fwd_bwd_us": 2329.7280073165894 + } + } + } + }, + "sglang": { + "diffusers[sglang]": { + "source": "diffusers AdaLN projection, BF16 F.linear [sglang]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 345.8240032196045, + "fwd_bwd_us": 2387.8400325775146 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 329.76000010967255, + "fwd_bwd_us": 3095.7919359207153 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 331.2000036239624, + "fwd_bwd_us": 3088.063955307007 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 0.002934168727885534, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 243.24800074100494, + "fwd_bwd_us": 3053.3759593963623 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.002284836389545336, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 311.16798520088196, + "fwd_bwd_us": 3689.5840167999268 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 1802.2559881210327, + "fwd_bwd_us": 14296.223640441895 + } + } + } + }, + "sglang_ieee": { + "diffusers[sglang_ieee]": { + "source": "diffusers AdaLN projection, BF16 F.linear [sglang_ieee]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 347.6160019636154, + "fwd_bwd_us": 2428.431987762451 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 332.0319950580597, + "fwd_bwd_us": 3137.839913368225 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 332.272008061409, + "fwd_bwd_us": 3133.8720321655273 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 0.002934168727885534, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 241.7599931359291, + "fwd_bwd_us": 3196.015954017639 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.002284836389545336, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 311.3119900226593, + "fwd_bwd_us": 3740.447998046875 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 1802.5919795036316, + "fwd_bwd_us": 14325.071811676025 + } + } + } + }, + "megatron_te_native": { + "diffusers[megatron_te_native]": { + "source": "diffusers AdaLN projection, BF16 F.linear [megatron_te_native]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 192 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 771 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 6144 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 2056 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 1104, + "first_failures": [ + { + "seed": 3, + "batch": 64, + "position": 0, + "row": 0 + }, + { + "seed": 3, + "batch": 64, + "position": 31, + "row": 0 + }, + { + "seed": 3, + "batch": 64, + "position": 63, + "row": 0 + }, + { + "seed": 3, + "batch": 64, + "position": 0, + "row": 259 + }, + { + "seed": 3, + "batch": 64, + "position": 31, + "row": 259 + }, + { + "seed": 3, + "batch": 64, + "position": 63, + "row": 259 + }, + { + "seed": 3, + "batch": 64, + "position": 0, + "row": 515 + }, + { + "seed": 3, + "batch": 64, + "position": 31, + "row": 515 + } + ] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 100.51199793815613, + "fwd_bwd_us": 759.440004825592 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 100.0640019774437, + "fwd_bwd_us": 702.351987361908 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 99.45599734783173, + "fwd_bwd_us": 687.5999867916107 + }, + "64": { + "fwd_rel_err": 0.003117907166490413, + "grad_rel_err": { + "temb": 0.0028456426244756147, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 111.16799712181091, + "fwd_bwd_us": 697.8560090065002 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.0022159011369222194, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 127.00800597667694, + "fwd_bwd_us": 706.032007932663 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 703.9839923381805, + "fwd_bwd_us": 2330.4959535598755 + } + } + } + }, + "megatron_triton": { + "diffusers[megatron_triton]": { + "source": "diffusers AdaLN projection, BF16 F.linear [megatron_triton]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 341.90399944782257, + "fwd_bwd_us": 2427.3279905319214 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 330.3360044956207, + "fwd_bwd_us": 2532.096028327942 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 329.23200726509094, + "fwd_bwd_us": 2525.215983390808 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 0.002934168727885534, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 236.40000075101852, + "fwd_bwd_us": 2594.655990600586 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.002284836389545336, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 305.00800907611847, + "fwd_bwd_us": 3204.032063484192 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 1738.0319833755493, + "fwd_bwd_us": 13771.424293518066 + } + } + } + }, + "megatron_triton_ieee": { + "diffusers[megatron_triton_ieee]": { + "source": "diffusers AdaLN projection, BF16 F.linear [megatron_triton_ieee]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.003856488128639106, + "grad_rel_err": { + "temb": 0.002906892643209935, + "w": 0.004108917853014856, + "b": 0.0 + }, + "fwd_us": 345.12001276016235, + "fwd_bwd_us": 2428.8480281829834 + }, + "3": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.0025072553507904586, + "w": 0.0039446027143119735, + "b": 0.0019723865877712033 + }, + "fwd_us": 329.8879861831665, + "fwd_bwd_us": 2532.6719284057617 + }, + "4": { + "fwd_rel_err": 0.0035999863805493103, + "grad_rel_err": { + "temb": 0.002560768077463324, + "w": 0.0030150633254733736, + "b": 0.0017331022530329288 + }, + "fwd_us": 329.0559947490692, + "fwd_bwd_us": 2532.0639610290527 + }, + "64": { + "fwd_rel_err": 0.0031179071664915094, + "grad_rel_err": { + "temb": 0.002934168727885534, + "w": 0.003786559211024591, + "b": 0.003089158155830867 + }, + "fwd_us": 237.42400109767914, + "fwd_bwd_us": 2599.3279218673706 + }, + "256": { + "fwd_rel_err": 0.0034422315551847975, + "grad_rel_err": { + "temb": 0.002284836389545336, + "w": 0.003184281504467339, + "b": 0.0025350368296686533 + }, + "fwd_us": 305.7119995355606, + "fwd_bwd_us": 3206.480026245117 + }, + "2048": { + "fwd_rel_err": 0.003219885589372577, + "grad_rel_err": { + "temb": 0.002307129848434764, + "w": 0.003537205254325983, + "b": 0.0020481905297133167 + }, + "fwd_us": 1738.5759949684143, + "fwd_bwd_us": 13772.655963897705 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/adaln_projection.png b/reports/experiments/h3-prior-art-b200/adaln_projection.png new file mode 100644 index 000000000..82188a66d Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/adaln_projection.png differ diff --git a/reports/experiments/h3-prior-art-b200/adaln_row_gather.json b/reports/experiments/h3-prior-art-b200/adaln_row_gather.json new file mode 100644 index 000000000..fea47a126 --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/adaln_row_gather.json @@ -0,0 +1,248 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "adaln_row_gather", + "rl_kernel_commit": "6c900ae0a7026230e0cdc905e55fc577b6f8d772", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": null, + "mode_environment": { + "plain": {} + } + }, + "quick": false, + "results": { + "plain": { + "diffusers": { + "source": "diffusers six index_select calls", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": false, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.0, + "grad_rel_err": { + "table": 0.06040967933388176 + }, + "fwd_us": 70.56000083684921, + "fwd_bwd_us": 1380.5919885635376 + }, + "32768": { + "fwd_rel_err": 0.0, + "grad_rel_err": { + "table": 0.2056162960711375 + }, + "fwd_us": 498.6880123615265, + "fwd_bwd_us": 9753.087997436523 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3AdaLNRowGatherCudaOp", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.0, + "grad_rel_err": { + "table": 0.0025124642229748073 + }, + "fwd_us": 115.11999741196632, + "fwd_bwd_us": 1105.679988861084 + }, + "32768": { + "fwd_rel_err": 0.0, + "grad_rel_err": { + "table": 0.00302351292151211 + }, + "fwd_us": 386.8959993124008, + "fwd_bwd_us": 2776.4800786972046 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/adaln_row_gather.png b/reports/experiments/h3-prior-art-b200/adaln_row_gather.png new file mode 100644 index 000000000..53542208f Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/adaln_row_gather.png differ diff --git a/reports/experiments/h3-prior-art-b200/final_adaln_out.json b/reports/experiments/h3-prior-art-b200/final_adaln_out.json new file mode 100644 index 000000000..bc33e6f51 --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/final_adaln_out.json @@ -0,0 +1,264 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "final_adaln_out", + "rl_kernel_commit": "03da729bc1b780ab1e5fd62f0dd654e82858dbed", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": null, + "mode_environment": { + "plain": {} + } + }, + "quick": false, + "results": { + "plain": { + "diffusers": { + "source": "diffusers MiniMaxH3AdaLayerNormOut (op-for-op replay)", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": false, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.008261982985661668, + "grad_rel_err": { + "x": 0.00678959940192795, + "nw": 0.0046959379804365755, + "temb": 0.03767563824438109, + "w": 0.07340738172976051, + "b": 0.04351708064268574 + }, + "fwd_us": 147.64800667762756, + "fwd_bwd_us": 1020.8479762077332 + }, + "32768": { + "fwd_rel_err": 0.008909737318797937, + "grad_rel_err": { + "x": 0.008348533984714676, + "nw": 0.00440358505956178, + "temb": 0.1248927861759312, + "w": 0.23805251261668708, + "b": 0.14489625509686074 + }, + "fwd_us": 779.6320021152496, + "fwd_bwd_us": 4525.4881381988525 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3FinalAdaLNOutCudaOp", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.008261982985661668, + "grad_rel_err": { + "x": 0.005084230038226309, + "nw": 0.00364253300425604, + "temb": 0.0010437274853488808, + "w": 0.00516260830853332, + "b": 0.0021386665111679487 + }, + "fwd_us": 131.4079985022545, + "fwd_bwd_us": 1641.5839791297913 + }, + "32768": { + "fwd_rel_err": 0.008909737318797937, + "grad_rel_err": { + "x": 0.005702424220873344, + "nw": 0.00440358505956178, + "temb": 0.0013565354049622875, + "w": 0.0038323877243695795, + "b": 0.002404863754273273 + }, + "fwd_us": 384.97599959373474, + "fwd_bwd_us": 2958.896040916443 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/final_adaln_out.png b/reports/experiments/h3-prior-art-b200/final_adaln_out.png new file mode 100644 index 000000000..cc0130f54 Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/final_adaln_out.png differ diff --git a/reports/experiments/h3-prior-art-b200/gate_residual.json b/reports/experiments/h3-prior-art-b200/gate_residual.json new file mode 100644 index 000000000..5339883ae --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/gate_residual.json @@ -0,0 +1,256 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "gate_residual", + "rl_kernel_commit": "401085437b33dfaac46f8fc4c593bc60bf04c2c9", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": null, + "mode_environment": { + "plain": {} + } + }, + "quick": false, + "results": { + "plain": { + "diffusers": { + "source": "diffusers residual + gate.index_select(...) * y", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": false, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.002994387720218618, + "grad_rel_err": { + "res": 0.0, + "y": 0.0025319460379000674, + "gate": 0.0426413874183677 + }, + "fwd_us": 62.04799935221672, + "fwd_bwd_us": 636.5279853343964 + }, + "32768": { + "fwd_rel_err": 0.0028753450414049685, + "grad_rel_err": { + "res": 0.0, + "y": 0.0025636917160711424, + "gate": 0.13599416946680765 + }, + "fwd_us": 410.38399934768677, + "fwd_bwd_us": 2502.079963684082 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3GateResidualCudaOp", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.002994387720218618, + "grad_rel_err": { + "res": 0.0, + "y": 0.0025319460379000674, + "gate": 0.0025585651639689682 + }, + "fwd_us": 79.45600152015686, + "fwd_bwd_us": 1210.256040096283 + }, + "32768": { + "fwd_rel_err": 0.0028753450414049685, + "grad_rel_err": { + "res": 0.0, + "y": 0.0025636917160711424, + "gate": 0.001969125163785725 + }, + "fwd_us": 279.79201078414917, + "fwd_bwd_us": 1595.6479907035828 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/gate_residual.png b/reports/experiments/h3-prior-art-b200/gate_residual.png new file mode 100644 index 000000000..cf923c9c7 Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/gate_residual.png differ diff --git a/reports/experiments/h3-prior-art-b200/norm_modulate.json b/reports/experiments/h3-prior-art-b200/norm_modulate.json new file mode 100644 index 000000000..b085baac9 --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/norm_modulate.json @@ -0,0 +1,791 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "norm_modulate", + "rl_kernel_commit": "ee83dec6bed474647ad11e99f038a3c02fd75183", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": null, + "mode_environment": { + "plain": {} + } + }, + "quick": false, + "results": { + "plain": { + "diffusers": { + "source": "diffusers composition (F.rms_norm + index_select modulation)", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": false, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.008581344842222902, + "grad_rel_err": { + "x": 0.005975029567735052, + "w": 0.003792743760107376, + "sh": 0.04506361745261099, + "sc": 0.035691318780065304 + }, + "fwd_us": 114.14399743080139, + "fwd_bwd_us": 872.2560107707977 + }, + "32768": { + "fwd_rel_err": 0.007507156057328251, + "grad_rel_err": { + "x": 0.005976764818120298, + "w": 0.004459331626195194, + "sh": 0.18642677708225106, + "sc": 0.09991998240326411 + }, + "fwd_us": 749.5839893817902, + "fwd_bwd_us": 4433.744192123413 + } + } + }, + "torch_rms_norm": { + "source": "torch F.rms_norm (no modulation)", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.0021347283487800277, + "grad_rel_err": { + "x": 0.002102632726644688, + "w": 0.0021263644360522546 + }, + "fwd_us": 25.119999423623085, + "fwd_bwd_us": 330.55999875068665 + }, + "32768": { + "fwd_rel_err": 0.0020067899681670957, + "grad_rel_err": { + "x": 0.002755700733757338, + "w": 0.002435435140412523 + }, + "fwd_us": 133.40799510478973, + "fwd_bwd_us": 850.2239882946014 + } + } + }, + "transformer_engine": { + "source": "TE 2.20.2 RMSNorm (no modulation)", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.0021347283487800277, + "grad_rel_err": { + "x": 0.002102632726644688 + }, + "fwd_us": 104.0160022675991, + "fwd_bwd_us": 470.880001783371 + }, + "32768": { + "fwd_rel_err": 0.0020067899681670957, + "grad_rel_err": { + "x": 0.002755700733757338 + }, + "fwd_us": 514.8159861564636, + "fwd_bwd_us": 1596.5759754180908 + } + } + }, + "liger": { + "source": "Liger 0.8.4 modulated RMSNorm + index_select", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": false, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.006623176661481121, + "grad_rel_err": { + "x": 0.006529533263925322, + "w": 0.004774428259970641, + "sh": 0.047092360830808686, + "sc": 0.04311728982829122 + }, + "fwd_us": 90.16000106930733, + "fwd_bwd_us": 809.8560273647308 + }, + "32768": { + "fwd_rel_err": 0.0068627233126289, + "grad_rel_err": { + "x": 0.008037730297978113, + "w": 0.00427513304978031, + "sh": 0.1954857309726789, + "sc": 0.10172031739215236 + }, + "fwd_us": 475.5519926548004, + "fwd_bwd_us": 3910.975933074951 + } + } + }, + "liger_rlk_gather": { + "source": "Liger 0.8.4 modulated RMSNorm + rl-kernel row gather", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.006623176661481121, + "grad_rel_err": { + "x": 0.006529533263925322, + "w": 0.004774428259970641, + "sh": 0.002599702677107761, + "sc": 0.0033097287013233684 + }, + "fwd_us": 185.29599905014038, + "fwd_bwd_us": 1824.1600394248962 + }, + "32768": { + "fwd_rel_err": 0.0068627233126289, + "grad_rel_err": { + "x": 0.008037730297978113, + "w": 0.00427513304978031, + "sh": 0.0018109614060863243, + "sc": 0.003963275077563095 + }, + "fwd_us": 692.9919719696045, + "fwd_bwd_us": 4762.704133987427 + } + } + }, + "sglang": { + "source": "SGLang 0.5.21 fused_norm_scale_shift (forward only)", + "has_backward": false, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.004400573087409222, + "fwd_us": 85.42400225996971 + }, + "32768": { + "fwd_rel_err": 0.0040637371636112985, + "fwd_us": 435.5680048465729 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3RMSNormCudaOp.forward_modulated", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 131072, + "sub_batch_sizes": [ + 1, + 7, + 4097, + 32768, + 65536 + ], + "sub_batches": 2124, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "4097": { + "fwd_rel_err": 0.008581344842222902, + "grad_rel_err": { + "x": 0.004236485456288316, + "w": 0.0030398817501858352, + "sh": 0.002599702677107761, + "sc": 0.002757974651813979 + }, + "fwd_us": 93.44000369310379, + "fwd_bwd_us": 1092.6080346107483 + }, + "32768": { + "fwd_rel_err": 0.007507156057328251, + "grad_rel_err": { + "x": 0.005048267823327812, + "w": 0.002926860714420976, + "sh": 0.0018109614060863243, + "sc": 0.003541578708222786 + }, + "fwd_us": 349.87199306488037, + "fwd_bwd_us": 2653.231978416443 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/norm_modulate.png b/reports/experiments/h3-prior-art-b200/norm_modulate.png new file mode 100644 index 000000000..6594edcd1 Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/norm_modulate.png differ diff --git a/reports/experiments/h3-prior-art-b200/timestep_mlp.json b/reports/experiments/h3-prior-art-b200/timestep_mlp.json new file mode 100644 index 000000000..38e1feb4c --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/timestep_mlp.json @@ -0,0 +1,1492 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "timestep_mlp", + "rl_kernel_commit": "ffd958afe27a081eee663552e7cc1d43dec8037b", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": "52fbcbc", + "mode_environment": { + "plain": {}, + "vllm": { + "CUBLAS_WORKSPACE_CONFIG": ":16:8", + "CUBLASLT_WORKSPACE_SIZE": "1" + }, + "sglang": {}, + "sglang_ieee": { + "TRITON_F32_DEFAULT": "ieee" + }, + "megatron_te_native": { + "CUBLASLT_WORKSPACE_SIZE": "0" + }, + "megatron_triton": {}, + "megatron_triton_ieee": { + "TRITON_F32_DEFAULT": "ieee" + } + } + }, + "quick": false, + "results": { + "plain": { + "diffusers[plain]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [plain]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 192, + "rowgrad_differ": 192 + }, + "257": { + "rows": 771, + "fwd_differ": 771, + "rowgrad_differ": 771 + }, + "2048": { + "rows": 6144, + "fwd_differ": 6144, + "rowgrad_differ": 6144 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 2048, + "rowgrad_differ": 2060 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 2232, + "first_failures": [ + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 771 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 771 + } + ] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.807769045764923e-07, + "grad_rel_err": { + "x": 3.2211968944911197e-07, + "w1": 1.5098534195344955e-07, + "b1": 1.410350051007604e-07, + "w2": 2.175078901511366e-07, + "b2": 0.0 + }, + "fwd_us": 36.927999928593636, + "fwd_bwd_us": 711.1679911613464 + }, + "3": { + "fwd_rel_err": 5.676673943201589e-07, + "grad_rel_err": { + "x": 4.085326731917541e-07, + "w1": 2.6733029568896263e-07, + "b1": 2.5557345343212385e-07, + "w2": 1.988937518622058e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 81.34400099515915, + "fwd_bwd_us": 776.8959999084473 + }, + "4": { + "fwd_rel_err": 5.676673937185161e-07, + "grad_rel_err": { + "x": 3.442394076935069e-07, + "w1": 2.901440780140815e-07, + "b1": 1.9831658889814138e-07, + "w2": 2.0747982241727246e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 81.55200257897377, + "fwd_bwd_us": 748.1759786605835 + }, + "64": { + "fwd_rel_err": 2.664792719694426e-06, + "grad_rel_err": { + "x": 4.3474725388849557e-07, + "w1": 4.887688242002194e-07, + "b1": 2.529915913033418e-07, + "w2": 5.072615732613012e-07, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 143.24799925088882, + "fwd_bwd_us": 756.7680180072784 + }, + "256": { + "fwd_rel_err": 5.712898230928172e-07, + "grad_rel_err": { + "x": 9.95098910262489e-07, + "w1": 9.672567531194054e-07, + "b1": 8.58617818657395e-07, + "w2": 7.957171915884723e-07, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 271.58400416374207, + "fwd_bwd_us": 832.3519825935364 + }, + "2048": { + "fwd_rel_err": 3.6510064550847246e-06, + "grad_rel_err": { + "x": 1.012898298588384e-06, + "w1": 1.1825830944307205e-06, + "b1": 8.824350021477035e-07, + "w2": 2.011824123714082e-06, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 1258.8639855384827, + "fwd_bwd_us": 3608.9760065078735 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3TimestepMLPCudaOp", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 2.421254715368641e-07, + "grad_rel_err": { + "x": 3.048727893223497e-07, + "w1": 1.9278019746544753e-07, + "b1": 1.7432439994844678e-07, + "w2": 1.9541577141227777e-07, + "b2": 0.0 + }, + "fwd_us": 35.88800132274628, + "fwd_bwd_us": 755.4719746112823 + }, + "3": { + "fwd_rel_err": 2.5720296794480815e-07, + "grad_rel_err": { + "x": 4.203330995116559e-07, + "w1": 2.6733029568896263e-07, + "b1": 2.5557345343212385e-07, + "w2": 1.685021626750072e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 44.544000178575516, + "fwd_bwd_us": 748.2079863548279 + }, + "4": { + "fwd_rel_err": 2.5720296775968735e-07, + "grad_rel_err": { + "x": 4.2277043456247276e-07, + "w1": 2.901440780140815e-07, + "b1": 2.1689270298354394e-07, + "w2": 1.8221364414193866e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 47.200001776218414, + "fwd_bwd_us": 737.1520102024078 + }, + "64": { + "fwd_rel_err": 2.7933997201982694e-07, + "grad_rel_err": { + "x": 6.024913562191266e-07, + "w1": 4.1809741087508027e-07, + "b1": 2.83039888633215e-07, + "w2": 4.358982508302425e-07, + "b2": 2.862777357761908e-07 + }, + "fwd_us": 448.2719898223877, + "fwd_bwd_us": 1322.0799565315247 + }, + "256": { + "fwd_rel_err": 2.2632120602315117e-07, + "grad_rel_err": { + "x": 5.171878010453915e-07, + "w1": 7.428217251157735e-07, + "b1": 4.2853189560371657e-07, + "w2": 7.631644366992322e-07, + "b2": 4.686685395440396e-07 + }, + "fwd_us": 1743.0559992790222, + "fwd_bwd_us": 4594.592094421387 + }, + "2048": { + "fwd_rel_err": 2.584209429418547e-07, + "grad_rel_err": { + "x": 4.99180005544531e-07, + "w1": 1.9541783623559206e-06, + "b1": 1.7851708196022288e-06, + "w2": 2.2796310999134285e-06, + "b2": 2.1701201580690495e-06 + }, + "fwd_us": 15600.304126739502, + "fwd_bwd_us": 42992.52891540527 + } + } + } + }, + "vllm": { + "diffusers[vllm]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [vllm]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 192, + "rowgrad_differ": 192 + }, + "257": { + "rows": 771, + "fwd_differ": 771, + "rowgrad_differ": 771 + }, + "2048": { + "rows": 6144, + "fwd_differ": 6144, + "rowgrad_differ": 6144 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 2048, + "rowgrad_differ": 2048 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 2232, + "first_failures": [ + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 771 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 771 + } + ] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.807769046831736e-07, + "grad_rel_err": { + "x": 2.6849865141614144e-07, + "w1": 1.5098534195344955e-07, + "b1": 1.410350051007604e-07, + "w2": 2.1750789001934508e-07, + "b2": 0.0 + }, + "fwd_us": 37.02399879693985, + "fwd_bwd_us": 728.3520102500916 + }, + "3": { + "fwd_rel_err": 1.7537237228724268e-07, + "grad_rel_err": { + "x": 4.139713215774079e-07, + "w1": 4.525846935277643e-07, + "b1": 3.8700315371982306e-07, + "w2": 1.988937517489883e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 55.63199892640114, + "fwd_bwd_us": 752.1600127220154 + }, + "4": { + "fwd_rel_err": 2.3380251126711334e-07, + "grad_rel_err": { + "x": 3.7918076562025883e-07, + "w1": 4.5526762841586136e-07, + "b1": 3.565643545954348e-07, + "w2": 2.0747982264348798e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 51.42400041222572, + "fwd_bwd_us": 744.8640167713165 + }, + "64": { + "fwd_rel_err": 2.664792717346965e-06, + "grad_rel_err": { + "x": 2.6721012462734666e-06, + "w1": 1.141325672600531e-06, + "b1": 8.36237373280394e-07, + "w2": 5.072615732613012e-07, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 165.40800034999847, + "fwd_bwd_us": 770.2080011367798 + }, + "256": { + "fwd_rel_err": 3.6510064555755425e-06, + "grad_rel_err": { + "x": 2.767490276754921e-06, + "w1": 9.67256753780727e-07, + "b1": 8.586178187169499e-07, + "w2": 7.957171915884723e-07, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 270.224004983902, + "fwd_bwd_us": 921.0879802703857 + }, + "2048": { + "fwd_rel_err": 3.6510064555755425e-06, + "grad_rel_err": { + "x": 3.1742638611753815e-06, + "w1": 2.089163185829833e-06, + "b1": 8.824350024581783e-07, + "w2": 2.011824123714082e-06, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 1285.3120565414429, + "fwd_bwd_us": 3825.792074203491 + } + } + } + }, + "sglang": { + "diffusers[sglang]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [sglang]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.0016499902522073425, + "grad_rel_err": { + "x": 0.0015402244992047199, + "w1": 0.0015883507645398632, + "b1": 0.0008914156583765084, + "w2": 0.0015720628608286314, + "b2": 0.0 + }, + "fwd_us": 159.93600338697433, + "fwd_bwd_us": 1034.015953540802 + }, + "3": { + "fwd_rel_err": 0.0015654570674600726, + "grad_rel_err": { + "x": 0.001624991665672256, + "w1": 0.001624925581907528, + "b1": 0.00096959190034644, + "w2": 0.0015932095611516203, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 161.72799468040466, + "fwd_bwd_us": 1035.5520248413086 + }, + "4": { + "fwd_rel_err": 0.0015828931805617757, + "grad_rel_err": { + "x": 0.0017129900427452491, + "w1": 0.0015986209418147807, + "b1": 0.0009882605763522905, + "w2": 0.0019152801810112664, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 163.85599970817566, + "fwd_bwd_us": 1044.3360209465027 + }, + "64": { + "fwd_rel_err": 0.0016120512469764673, + "grad_rel_err": { + "x": 0.0016110202455347832, + "w1": 0.0017370322406895582, + "b1": 0.0008156819272919207, + "w2": 0.0018132309973686246, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 140.27199894189835, + "fwd_bwd_us": 1058.8319897651672 + }, + "256": { + "fwd_rel_err": 0.0014704288089542434, + "grad_rel_err": { + "x": 0.0015138544453058315, + "w1": 0.0015415965393197001, + "b1": 0.0009163867703221417, + "w2": 0.0015096106678851324, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 154.38400208950043, + "fwd_bwd_us": 1140.2400135993958 + }, + "2048": { + "fwd_rel_err": 0.0014704288089542434, + "grad_rel_err": { + "x": 0.001664412754369862, + "w1": 0.0015475258170343411, + "b1": 0.0008733500142339925, + "w2": 0.0016480881496351136, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 349.5360016822815, + "fwd_bwd_us": 2636.399984359741 + } + } + } + }, + "sglang_ieee": { + "diffusers[sglang_ieee]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [sglang_ieee]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 2.4718486470770298e-06, + "grad_rel_err": { + "x": 2.2013342726856412e-06, + "w1": 1.287841337898046e-06, + "b1": 1.2723330547272893e-06, + "w2": 6.483615007349123e-07, + "b2": 0.0 + }, + "fwd_us": 1359.0400218963623, + "fwd_bwd_us": 2278.7680625915527 + }, + "3": { + "fwd_rel_err": 3.257253677979414e-06, + "grad_rel_err": { + "x": 2.3470823967651816e-06, + "w1": 1.3245040616204118e-06, + "b1": 1.2442139377924535e-06, + "w2": 5.381404258436617e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 1330.2080035209656, + "fwd_bwd_us": 2234.0800762176514 + }, + "4": { + "fwd_rel_err": 3.2572536778868546e-06, + "grad_rel_err": { + "x": 2.149831291199852e-06, + "w1": 1.5529519087445555e-06, + "b1": 1.188753149346917e-06, + "w2": 5.637581194862654e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 1333.952009677887, + "fwd_bwd_us": 2233.296036720276 + }, + "64": { + "fwd_rel_err": 2.664792719694426e-06, + "grad_rel_err": { + "x": 2.672101244999492e-06, + "w1": 1.1413256730557414e-06, + "b1": 8.362373736199796e-07, + "w2": 5.072615732613012e-07, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 1361.9359731674194, + "fwd_bwd_us": 2299.3119955062866 + }, + "256": { + "fwd_rel_err": 3.6510064550847246e-06, + "grad_rel_err": { + "x": 2.767490278047356e-06, + "w1": 9.672567531194054e-07, + "b1": 8.58617818657395e-07, + "w2": 7.957171915884723e-07, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 1364.3839955329895, + "fwd_bwd_us": 2428.3519983291626 + }, + "2048": { + "fwd_rel_err": 3.6510064550847246e-06, + "grad_rel_err": { + "x": 3.174263860695608e-06, + "w1": 2.0891631849115376e-06, + "b1": 8.824350021477035e-07, + "w2": 2.011824123714082e-06, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 4205.615997314453, + "fwd_bwd_us": 7328.56011390686 + } + } + } + }, + "megatron_te_native": { + "diffusers[megatron_te_native]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [megatron_te_native]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 192, + "rowgrad_differ": 192 + }, + "257": { + "rows": 771, + "fwd_differ": 771, + "rowgrad_differ": 771 + }, + "2048": { + "rows": 6144, + "fwd_differ": 6144, + "rowgrad_differ": 6144 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 2048, + "rowgrad_differ": 2060 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 2232, + "first_failures": [ + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 0 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 259 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 515 + }, + { + "seed": 3, + "batch": 2, + "position": 0, + "row": 771 + }, + { + "seed": 3, + "batch": 2, + "position": 1, + "row": 771 + } + ] + }, + "batch_invariant": false + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.807769046831736e-07, + "grad_rel_err": { + "x": 3.2211968936309937e-07, + "w1": 1.5098534195344955e-07, + "b1": 1.410350051007604e-07, + "w2": 2.1750789001934508e-07, + "b2": 0.0 + }, + "fwd_us": 36.81600093841553, + "fwd_bwd_us": 700.4320025444031 + }, + "3": { + "fwd_rel_err": 1.7537237228724268e-07, + "grad_rel_err": { + "x": 4.0853267297182194e-07, + "w1": 2.6733029568896263e-07, + "b1": 2.5557345349749303e-07, + "w2": 1.988937517489883e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 55.75999990105629, + "fwd_bwd_us": 729.4560074806213 + }, + "4": { + "fwd_rel_err": 2.3380251126711334e-07, + "grad_rel_err": { + "x": 3.442394076935068e-07, + "w1": 2.9014407801408155e-07, + "b1": 1.9831658878698316e-07, + "w2": 2.0747982264348798e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 49.6320016682148, + "fwd_bwd_us": 730.5920124053955 + }, + "64": { + "fwd_rel_err": 2.664792717346965e-06, + "grad_rel_err": { + "x": 4.3474725388849557e-07, + "w1": 4.887688242002194e-07, + "b1": 2.529915913033418e-07, + "w2": 5.072615732613012e-07, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 150.751993060112, + "fwd_bwd_us": 748.6079931259155 + }, + "256": { + "fwd_rel_err": 3.6510064555755425e-06, + "grad_rel_err": { + "x": 9.95098910262489e-07, + "w1": 9.672567531194054e-07, + "b1": 8.58617818657395e-07, + "w2": 7.957171915884723e-07, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 270.9920108318329, + "fwd_bwd_us": 814.4640028476715 + }, + "2048": { + "fwd_rel_err": 3.6510064555755425e-06, + "grad_rel_err": { + "x": 1.012898298588384e-06, + "w1": 1.1825830944307205e-06, + "b1": 8.824350021477035e-07, + "w2": 2.011824123714082e-06, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 1285.6639623641968, + "fwd_bwd_us": 3635.6799602508545 + } + } + } + }, + "megatron_triton": { + "diffusers[megatron_triton]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [megatron_triton]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 0.0016499902522073425, + "grad_rel_err": { + "x": 0.0015402244992047199, + "w1": 0.0015883507645398632, + "b1": 0.0008914156583765084, + "w2": 0.0015720628608286314, + "b2": 0.0 + }, + "fwd_us": 157.51999616622925, + "fwd_bwd_us": 1031.4080119132996 + }, + "3": { + "fwd_rel_err": 0.0015654570674600726, + "grad_rel_err": { + "x": 0.001624991665672256, + "w1": 0.001624925581907528, + "b1": 0.00096959190034644, + "w2": 0.0015932095611516203, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 158.38400274515152, + "fwd_bwd_us": 1026.5119671821594 + }, + "4": { + "fwd_rel_err": 0.0015828931805617757, + "grad_rel_err": { + "x": 0.0017129900427452491, + "w1": 0.0015986209418147807, + "b1": 0.0009882605763522905, + "w2": 0.0019152801810112664, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 160.3040024638176, + "fwd_bwd_us": 1025.2799987792969 + }, + "64": { + "fwd_rel_err": 0.0016120512469764673, + "grad_rel_err": { + "x": 0.0016110202455347832, + "w1": 0.0017370322406895582, + "b1": 0.0008156819272919207, + "w2": 0.0018132309973686246, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 137.35999912023544, + "fwd_bwd_us": 1050.607979297638 + }, + "256": { + "fwd_rel_err": 0.0014704288089542434, + "grad_rel_err": { + "x": 0.0015138544453058315, + "w1": 0.0015415965393197001, + "b1": 0.0009163867703221417, + "w2": 0.0015096106678851324, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 151.47200226783752, + "fwd_bwd_us": 1136.73597574234 + }, + "2048": { + "fwd_rel_err": 0.0014704288089542434, + "grad_rel_err": { + "x": 0.001664412754369862, + "w1": 0.0015475258170343411, + "b1": 0.0008733500142339925, + "w2": 0.0016480881496351136, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 348.9600121974945, + "fwd_bwd_us": 2634.1439485549927 + } + } + } + }, + "megatron_triton_ieee": { + "diffusers[megatron_triton_ieee]": { + "source": "diffusers TimestepEmbedding, FP32 F.linear [megatron_triton_ieee]", + "has_backward": true, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": true, + "full_vs_sub_batches": { + "full_batch": 4096, + "sub_batch_sizes": [ + 1, + 7, + 1024, + 2048 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 2.4718486470770298e-06, + "grad_rel_err": { + "x": 2.2013342726856412e-06, + "w1": 1.287841337898046e-06, + "b1": 1.2723330547272893e-06, + "w2": 6.483615007349123e-07, + "b2": 0.0 + }, + "fwd_us": 1357.26398229599, + "fwd_bwd_us": 2277.535915374756 + }, + "3": { + "fwd_rel_err": 3.257253677979414e-06, + "grad_rel_err": { + "x": 2.3470823967651816e-06, + "w1": 1.3245040616204118e-06, + "b1": 1.2442139377924535e-06, + "w2": 5.381404258436617e-07, + "b2": 4.7578267757105583e-08 + }, + "fwd_us": 1399.4399905204773, + "fwd_bwd_us": 2308.751940727234 + }, + "4": { + "fwd_rel_err": 3.2572536778868546e-06, + "grad_rel_err": { + "x": 2.149831291199852e-06, + "w1": 1.5529519087445555e-06, + "b1": 1.188753149346917e-06, + "w2": 5.637581194862654e-07, + "b2": 6.404684939390106e-08 + }, + "fwd_us": 1400.4319906234741, + "fwd_bwd_us": 2306.9440126419067 + }, + "64": { + "fwd_rel_err": 2.664792719694426e-06, + "grad_rel_err": { + "x": 2.672101244999492e-06, + "w1": 1.1413256730557414e-06, + "b1": 8.362373736199796e-07, + "w2": 5.072615732613012e-07, + "b2": 8.671140993863786e-08 + }, + "fwd_us": 1377.2000074386597, + "fwd_bwd_us": 2302.672028541565 + }, + "256": { + "fwd_rel_err": 3.6510064550847246e-06, + "grad_rel_err": { + "x": 2.767490278047356e-06, + "w1": 9.672567531194054e-07, + "b1": 8.58617818657395e-07, + "w2": 7.957171915884723e-07, + "b2": 1.2705827741456392e-07 + }, + "fwd_us": 1380.2080154418945, + "fwd_bwd_us": 2435.231924057007 + }, + "2048": { + "fwd_rel_err": 3.6510064550847246e-06, + "grad_rel_err": { + "x": 3.174263860695608e-06, + "w1": 2.0891631849115376e-06, + "b1": 8.824350021477035e-07, + "w2": 2.011824123714082e-06, + "b2": 1.0314665338164682e-07 + }, + "fwd_us": 4250.479936599731, + "fwd_bwd_us": 7348.78396987915 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/timestep_mlp.png b/reports/experiments/h3-prior-art-b200/timestep_mlp.png new file mode 100644 index 000000000..e1c285ef4 Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/timestep_mlp.png differ diff --git a/reports/experiments/h3-prior-art-b200/timestep_sinusoid.json b/reports/experiments/h3-prior-art-b200/timestep_sinusoid.json new file mode 100644 index 000000000..8cbea742b --- /dev/null +++ b/reports/experiments/h3-prior-art-b200/timestep_sinusoid.json @@ -0,0 +1,374 @@ +{ + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": "timestep_sinusoid", + "rl_kernel_commit": "7a82917c270cb42bc3350c1d30f97919b954dd86", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "libraries": { + "torch": "2.13.0", + "triton": "3.7.1", + "diffusers": null, + "vllm": "0.30.0", + "flashinfer-python": "0.6.18.post1", + "sglang": "0.5.21", + "liger-kernel": "0.8.4", + "transformer_engine": "2.20.2", + "megatron-core": null + }, + "megatron_source_commit": null, + "mode_environment": { + "plain": {} + } + }, + "quick": false, + "results": { + "plain": { + "diffusers": { + "source": "diffusers get_timestep_embedding (op-for-op replay)", + "has_backward": false, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 8192, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.205756250743063e-08, + "fwd_us": 63.58399987220764 + }, + "3": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 63.87200206518173 + }, + "4": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 63.920002430677414 + }, + "64": { + "fwd_rel_err": 7.54632437649728e-08, + "fwd_us": 64.01600316166878 + }, + "256": { + "fwd_rel_err": 8.564062026206226e-08, + "fwd_us": 64.2399974167347 + }, + "2048": { + "fwd_rel_err": 9.948448131957869e-08, + "fwd_us": 64.55999985337257 + } + } + }, + "sglang": { + "source": "SGLang 0.5.21 timestep_embedding", + "has_backward": false, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 8192, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.205756250743063e-08, + "fwd_us": 12.512000277638435 + }, + "3": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 12.512000277638435 + }, + "4": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 12.608000077307224 + }, + "64": { + "fwd_rel_err": 7.54632437649728e-08, + "fwd_us": 12.559999711811543 + }, + "256": { + "fwd_rel_err": 8.564062026206226e-08, + "fwd_us": 12.480000033974648 + }, + "2048": { + "fwd_rel_err": 9.948448131957869e-08, + "fwd_us": 12.543999589979649 + } + } + }, + "rl_kernel": { + "source": "rl-kernel H3TimestepSinusoidCudaOp (check_range=False)", + "has_backward": false, + "batch_invariance": { + "all_rows": { + "64": { + "rows": 192, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "257": { + "rows": 771, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "2048": { + "rows": 6144, + "fwd_differ": 0, + "rowgrad_differ": 0 + } + }, + "params_repeatable": null, + "full_vs_sub_batches": { + "full_batch": 8192, + "sub_batch_sizes": [ + 1, + 7, + 2048, + 4097 + ], + "sub_batches": 2060, + "fwd_differ": 0, + "rowgrad_differ": 0 + }, + "batch_size_sweep": { + "batch_sizes": [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 15, + 16, + 17, + 31, + 32, + 33, + 63, + 64, + 65, + 127, + 128, + 129, + 255, + 256, + 257, + 511, + 512, + 513, + 1023, + 1024, + 1025, + 2047, + 2048 + ], + "probe_rows": [ + 0, + 259, + 515, + 771, + 1027, + 1283, + 1539, + 2047 + ], + "comparisons": 2256, + "differ": 0, + "first_failures": [] + }, + "batch_invariant": true + }, + "accuracy_latency": { + "1": { + "fwd_rel_err": 3.205756250743063e-08, + "fwd_us": 15.48799965530634 + }, + "3": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 15.503999777138233 + }, + "4": { + "fwd_rel_err": 6.70917436344131e-08, + "fwd_us": 15.519999898970127 + }, + "64": { + "fwd_rel_err": 7.54632437649728e-08, + "fwd_us": 15.48799965530634 + }, + "256": { + "fwd_rel_err": 8.564062026206226e-08, + "fwd_us": 15.536000020802021 + }, + "2048": { + "fwd_rel_err": 9.948448131957869e-08, + "fwd_us": 15.552000142633915 + } + } + } + } + } +} diff --git a/reports/experiments/h3-prior-art-b200/timestep_sinusoid.png b/reports/experiments/h3-prior-art-b200/timestep_sinusoid.png new file mode 100644 index 000000000..a9815f486 Binary files /dev/null and b/reports/experiments/h3-prior-art-b200/timestep_sinusoid.png differ diff --git a/reports/experiments/h3-rmsnorm-b200/figure.png b/reports/experiments/h3-rmsnorm-b200/figure.png new file mode 100644 index 000000000..75c77fb07 Binary files /dev/null and b/reports/experiments/h3-rmsnorm-b200/figure.png differ diff --git a/reports/experiments/h3-rmsnorm-b200/report.json b/reports/experiments/h3-rmsnorm-b200/report.json new file mode 100644 index 000000000..dc6046ffb --- /dev/null +++ b/reports/experiments/h3-rmsnorm-b200/report.json @@ -0,0 +1,154 @@ +{ + "kind": "h3_operator_report", + "op": "h3_rmsnorm", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "80e46097302132c633dc6fc1a311097b616ac966", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "plain_bitwise_vs_nn_rmsnorm": { + "transformer_blocks.0.norm1.weight": true, + "transformer_blocks.0.norm2.weight": true, + "token_refiner.final_norm.weight": true, + "norm_out.norm.weight": true + }, + "modulated_bitwise_vs_diffusers": true, + "rows_batch_invariant": true, + "backward": { + "cuda": { + "repeat_bitwise_equal": true, + "rel_error": { + "dx": 0.004190817029101904, + "dweight": 0.0037932313344362263, + "dshift": 0.00256468183980526, + "dscale": 0.0028309973965559895 + } + }, + "provider": { + "repeat_bitwise_equal": false, + "rel_error": { + "dx": 0.006362059284075666, + "dweight": 0.0037932313344362263, + "dshift": 0.055593107435841636, + "dscale": 0.046512377475823756 + } + } + } + }, + "perf": [ + { + "op": "h3_rmsnorm", + "case": "S=4097", + "backend": "H3RMSNormCudaOp", + "bytes": 88101888, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 80.78400045633316, + "candidate_gbps": 1090.5858524253506, + "candidate_peak_mib": 42.0263671875, + "provider_us": 120.7519993185997, + "provider_gbps": 729.6101803461358, + "provider_peak_mib": 168.041015625, + "candidate_backward_us": 1229.6479940414429, + "candidate_backward_gbps": 71.64805572563778, + "candidate_backward_peak_mib": 43.54150390625, + "provider_backward_us": 849.2799997329712, + "provider_backward_gbps": 103.73715150209685, + "provider_backward_peak_mib": 126.123046875 + }, + { + "op": "h3_rmsnorm", + "case": "S=32768", + "backend": "H3RMSNormCudaOp", + "bytes": 704643072, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 345.71200609207153, + "candidate_gbps": 2038.2372020147207, + "candidate_peak_mib": 336.125, + "provider_us": 760.4160010814667, + "provider_gbps": 926.654713995831, + "provider_peak_mib": 1344.0, + "candidate_backward_us": 2180.527925491333, + "candidate_backward_gbps": 323.1525098864416, + "candidate_backward_peak_mib": 345.00390625, + "provider_backward_us": 4404.2558670043945, + "provider_backward_gbps": 159.9914022432287, + "provider_backward_peak_mib": 1008.09228515625 + }, + { + "op": "h3_rmsnorm", + "case": "S=131072", + "backend": "H3RMSNormCudaOp", + "bytes": 2818572288, + "execution_order": { + "policy": "alternating", + "iteration_0": [ + "candidate", + "provider", + "candidate_backward", + "provider_backward" + ], + "iteration_1": [ + "provider_backward", + "candidate_backward", + "provider", + "candidate" + ] + }, + "backward_timing_scope": "backward_only", + "candidate_us": 1209.1839909553528, + "candidate_gbps": 2330.9705628612405, + "candidate_peak_mib": 1344.5, + "provider_us": 2901.5519618988037, + "provider_gbps": 971.4016240314024, + "provider_peak_mib": 5376.0, + "candidate_backward_us": 6172.671794891357, + "candidate_backward_gbps": 456.62111669904664, + "candidate_backward_peak_mib": 1377.970703125, + "provider_backward_us": 17553.45630645752, + "provider_backward_gbps": 160.57078667539176, + "provider_backward_peak_mib": 4032.09228515625 + } + ] +} diff --git a/reports/experiments/h3-sp-norm-adaln-b200/figure.png b/reports/experiments/h3-sp-norm-adaln-b200/figure.png new file mode 100644 index 000000000..cae82445b Binary files /dev/null and b/reports/experiments/h3-sp-norm-adaln-b200/figure.png differ diff --git a/reports/experiments/h3-sp-norm-adaln-b200/report.json b/reports/experiments/h3-sp-norm-adaln-b200/report.json new file mode 100644 index 000000000..0cb326dab --- /dev/null +++ b/reports/experiments/h3-sp-norm-adaln-b200/report.json @@ -0,0 +1,1750 @@ +{ + "kind": "h3_ws2_report", + "op": "sp_norm_adaln", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "rl_kernel_commit": "6c0d3805ed7606fb120717de95c416aa597bc479", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "gpus": 8 + }, + "region": "norm2(residual + gate_msa[row] * y) * (1 + scale_mlp[row]) + shift_mlp[row]", + "hidden": 5376, + "runs": [ + { + "sp": 1, + "ranks": [ + { + "rank": 0, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4097, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 991.9680058956146, + "ws1_forward_backward_us": 2041.6001081466675 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4097, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 1059.2960119247437, + "ws1_forward_backward_us": 2027.567982673645 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8194, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 1368.9599633216858, + "ws1_forward_backward_us": 1929.3439984321594 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 3118.7199354171753, + "ws1_forward_backward_us": 3721.4879989624023 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 3082.8800201416016, + "ws1_forward_backward_us": 3720.639944076538 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 131072, + "rows_sent": 0, + "rows_gathered_per_rank": 0, + "sp_forward_backward_us": 10341.423988342285, + "ws1_forward_backward_us": 10732.816219329834 + } + ] + } + ] + }, + { + "sp": 2, + "ranks": [ + { + "rank": 0, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2048, + "rows_sent": 0, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1718.991994857788, + "ws1_forward_backward_us": 2053.599953651428 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2048, + "rows_sent": 0, + "rows_gathered_per_rank": 256, + "sp_forward_backward_us": 1728.9119958877563, + "ws1_forward_backward_us": 2043.280005455017 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 255, + "rows_gathered_per_rank": 311, + "sp_forward_backward_us": 1925.8559942245483, + "ws1_forward_backward_us": 1929.199993610382 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 0, + "rows_gathered_per_rank": 109, + "sp_forward_backward_us": 2848.639965057373, + "ws1_forward_backward_us": 3724.33602809906 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 0, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3447.7440118789673, + "ws1_forward_backward_us": 3744.6560859680176 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 65536, + "rows_sent": 0, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 7267.855882644653, + "ws1_forward_backward_us": 10750.207901000977 + } + ] + }, + { + "rank": 1, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2049, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1730.4319739341736, + "ws1_forward_backward_us": 2045.8720922470093 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2049, + "rows_sent": 256, + "rows_gathered_per_rank": 256, + "sp_forward_backward_us": 1732.5920462608337, + "ws1_forward_backward_us": 2017.7119970321655 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4098, + "rows_sent": 311, + "rows_gathered_per_rank": 311, + "sp_forward_backward_us": 1922.8639602661133, + "ws1_forward_backward_us": 2196.8480348587036 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 109, + "rows_gathered_per_rank": 109, + "sp_forward_backward_us": 2852.7839183807373, + "ws1_forward_backward_us": 3739.6160364151 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 1280, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3444.591999053955, + "ws1_forward_backward_us": 3704.927921295166 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 65536, + "rows_sent": 137, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 7268.079996109009, + "ws1_forward_backward_us": 10725.9840965271 + } + ] + } + ] + }, + { + "sp": 4, + "ranks": [ + { + "rank": 0, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 0, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1881.7440271377563, + "ws1_forward_backward_us": 2071.712017059326 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 0, + "rows_gathered_per_rank": 1025, + "sp_forward_backward_us": 3123.103976249695, + "ws1_forward_backward_us": 2100.256085395813 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2048, + "rows_sent": 255, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 2352.3999452590942, + "ws1_forward_backward_us": 2243.775963783264 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 0, + "rows_gathered_per_rank": 148, + "sp_forward_backward_us": 2398.591995239258, + "ws1_forward_backward_us": 3750.032067298889 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 0, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3881.5521001815796, + "ws1_forward_backward_us": 3646.672010421753 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 0, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 4657.984018325806, + "ws1_forward_backward_us": 10765.82384109497 + } + ] + }, + { + "rank": 1, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 38, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1882.4959993362427, + "ws1_forward_backward_us": 2080.8639526367188 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 1024, + "rows_gathered_per_rank": 1025, + "sp_forward_backward_us": 3116.8479919433594, + "ws1_forward_backward_us": 2126.1119842529297 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2048, + "rows_sent": 349, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 2355.2639484405518, + "ws1_forward_backward_us": 2273.632049560547 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 148, + "rows_gathered_per_rank": 148, + "sp_forward_backward_us": 2385.3759765625, + "ws1_forward_backward_us": 3814.9280548095703 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 1024, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3878.3520460128784, + "ws1_forward_backward_us": 3743.328094482422 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 119, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 4659.152030944824, + "ws1_forward_backward_us": 10744.271755218506 + } + ] + }, + { + "rank": 2, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1882.1280002593994, + "ws1_forward_backward_us": 2079.1680812835693 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 256, + "rows_gathered_per_rank": 1025, + "sp_forward_backward_us": 3128.880023956299, + "ws1_forward_backward_us": 2028.91206741333 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2048, + "rows_sent": 470, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 2358.7679862976074, + "ws1_forward_backward_us": 2218.4159755706787 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 109, + "rows_gathered_per_rank": 148, + "sp_forward_backward_us": 2398.159980773926, + "ws1_forward_backward_us": 3736.575961112976 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 1280, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3878.399968147278, + "ws1_forward_backward_us": 3727.1039485931396 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 137, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 4655.616044998169, + "ws1_forward_backward_us": 10773.631572723389 + } + ] + }, + { + "rank": 3, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1025, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 1882.3360204696655, + "ws1_forward_backward_us": 2083.999991416931 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1025, + "rows_sent": 1025, + "rows_gathered_per_rank": 1025, + "sp_forward_backward_us": 3124.8480081558228, + "ws1_forward_backward_us": 2027.2639989852905 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 2050, + "rows_sent": 444, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 2356.752038002014, + "ws1_forward_backward_us": 2231.839895248413 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 109, + "rows_gathered_per_rank": 148, + "sp_forward_backward_us": 2393.328070640564, + "ws1_forward_backward_us": 3794.848084449768 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 8192, + "rows_sent": 768, + "rows_gathered_per_rank": 1280, + "sp_forward_backward_us": 3877.5999546051025, + "ws1_forward_backward_us": 3747.183918952942 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 32768, + "rows_sent": 137, + "rows_gathered_per_rank": 137, + "sp_forward_backward_us": 4659.968137741089, + "ws1_forward_backward_us": 10759.23204421997 + } + ] + } + ] + }, + { + "sp": 8, + "ranks": [ + { + "rank": 0, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 0, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2426.2399673461914, + "ws1_forward_backward_us": 2074.0320682525635 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 0, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3422.368049621582, + "ws1_forward_backward_us": 2055.8719635009766 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 255, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3257.1520805358887, + "ws1_forward_backward_us": 2252.463936805725 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 0, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2719.6160554885864, + "ws1_forward_backward_us": 3762.287974357605 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 0, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6810.816049575806, + "ws1_forward_backward_us": 3834.04803276062 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 0, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3908.128023147583, + "ws1_forward_backward_us": 10828.735828399658 + } + ] + }, + { + "rank": 1, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 53, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2427.807927131653, + "ws1_forward_backward_us": 2092.944025993347 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 512, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3423.2640266418457, + "ws1_forward_backward_us": 2113.968014717102 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 349, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3257.823944091797, + "ws1_forward_backward_us": 2243.280053138733 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 148, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2718.0320024490356, + "ws1_forward_backward_us": 3741.8400049209595 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 512, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6815.024137496948, + "ws1_forward_backward_us": 3744.8959350585938 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 119, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3922.8639602661133, + "ws1_forward_backward_us": 10726.768016815186 + } + ] + }, + { + "rank": 2, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 38, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2428.831934928894, + "ws1_forward_backward_us": 2075.711965560913 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 512, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3417.6799058914185, + "ws1_forward_backward_us": 2049.488067626953 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 349, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3250.0, + "ws1_forward_backward_us": 2330.128073692322 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 148, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2713.855981826782, + "ws1_forward_backward_us": 3741.7439222335815 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 1024, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6805.504083633423, + "ws1_forward_backward_us": 3743.4879541397095 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 119, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3919.167995452881, + "ws1_forward_backward_us": 10753.407955169678 + } + ] + }, + { + "rank": 3, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 14, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2423.8399267196655, + "ws1_forward_backward_us": 2082.304000854492 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 512, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3420.400023460388, + "ws1_forward_backward_us": 2103.7439107894897 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 334, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3251.9359588623047, + "ws1_forward_backward_us": 1981.935977935791 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 224, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2718.2559967041016, + "ws1_forward_backward_us": 3831.8400382995605 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 1536, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6809.2639446258545, + "ws1_forward_backward_us": 3867.7600622177124 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 197, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3919.6159839630127, + "ws1_forward_backward_us": 10731.248378753662 + } + ] + }, + { + "rank": 4, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2424.607992172241, + "ws1_forward_backward_us": 2085.9040021896362 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 256, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3420.431971549988, + "ws1_forward_backward_us": 2131.3600540161133 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 470, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3255.4880380630493, + "ws1_forward_backward_us": 2287.984013557434 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 109, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2717.9359197616577, + "ws1_forward_backward_us": 3793.8079833984375 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 1280, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6811.775922775269, + "ws1_forward_backward_us": 3852.784037590027 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 137, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3918.784022331238, + "ws1_forward_backward_us": 10732.5758934021 + } + ] + }, + { + "rank": 5, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2418.5279607772827, + "ws1_forward_backward_us": 2071.8079805374146 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 512, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3410.7840061187744, + "ws1_forward_backward_us": 2094.015955924988 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 444, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3253.4559965133667, + "ws1_forward_backward_us": 2252.384066581726 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 109, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2717.3759937286377, + "ws1_forward_backward_us": 3766.095995903015 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 768, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6812.5598430633545, + "ws1_forward_backward_us": 3759.584069252014 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 137, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3901.247978210449, + "ws1_forward_backward_us": 10712.016105651855 + } + ] + }, + { + "rank": 6, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 223, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2416.319966316223, + "ws1_forward_backward_us": 2091.7919874191284 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 512, + "rows_sent": 512, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3417.8720712661743, + "ws1_forward_backward_us": 2062.7999305725098 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1024, + "rows_sent": 444, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3253.648042678833, + "ws1_forward_backward_us": 2388.256072998047 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 109, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2718.9279794692993, + "ws1_forward_backward_us": 3857.6799631118774 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 768, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6809.920072555542, + "ws1_forward_backward_us": 3778.880000114441 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 137, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3904.7359228134155, + "ws1_forward_backward_us": 10760.511875152588 + } + ] + }, + { + "rank": 7, + "cases": [ + { + "batch": 1, + "seq": 4097, + "layout": "block", + "seed": 1, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 513, + "rows_sent": 144, + "rows_gathered_per_rank": 223, + "sp_forward_backward_us": 2423.535943031311, + "ws1_forward_backward_us": 2088.44792842865 + }, + { + "batch": 1, + "seq": 4097, + "layout": "interleaved", + "seed": 2, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 513, + "rows_sent": 513, + "rows_gathered_per_rank": 513, + "sp_forward_backward_us": 3406.9759845733643, + "ws1_forward_backward_us": 2153.6799669265747 + }, + { + "batch": 2, + "seq": 4097, + "layout": "block", + "seed": 3, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 1026, + "rows_sent": 444, + "rows_gathered_per_rank": 470, + "sp_forward_backward_us": 3250.111937522888, + "ws1_forward_backward_us": 2282.0160388946533 + }, + { + "batch": 1, + "seq": 32768, + "layout": "block", + "seed": 4, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 109, + "rows_gathered_per_rank": 224, + "sp_forward_backward_us": 2721.60005569458, + "ws1_forward_backward_us": 3737.679958343506 + }, + { + "batch": 1, + "seq": 32768, + "layout": "interleaved", + "seed": 5, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 4096, + "rows_sent": 1280, + "rows_gathered_per_rank": 1536, + "sp_forward_backward_us": 6815.423965454102, + "ws1_forward_backward_us": 3777.9200077056885 + }, + { + "batch": 1, + "seq": 131072, + "layout": "block", + "seed": 6, + "equal": { + "out": true, + "d_residual": true, + "d_y": true, + "d_weight": true, + "d_table": true + }, + "local_rows": 16384, + "rows_sent": 137, + "rows_gathered_per_rank": 197, + "sp_forward_backward_us": 3918.0320501327515, + "ws1_forward_backward_us": 10712.448120117188 + } + ] + } + ] + } + ], + "naive_sp": { + "sp": 8, + "block": { + "d_weight": 0.4032738208770752, + "d_table": 0.06425677984952927 + }, + "interleaved": { + "d_weight": 0.4107142984867096, + "d_table": 0.20648698508739471 + } + } +} diff --git a/reports/experiments/h3-timestep-mlp-b200/figure.png b/reports/experiments/h3-timestep-mlp-b200/figure.png new file mode 100644 index 000000000..99e423a87 Binary files /dev/null and b/reports/experiments/h3-timestep-mlp-b200/figure.png differ diff --git a/reports/experiments/h3-timestep-mlp-b200/report.json b/reports/experiments/h3-timestep-mlp-b200/report.json new file mode 100644 index 000000000..7ebf89e2d --- /dev/null +++ b/reports/experiments/h3-timestep-mlp-b200/report.json @@ -0,0 +1,480 @@ +{ + "kind": "h3_operator_report", + "op": "timestep_mlp_fp32", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "65ef7f605c165a87be9271e09d19b0dbd8f5a4f6", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "contract_atol": 0.0001, + "draws": 200, + "num_timesteps": 4, + "cuda_max_abs_vs_fp64": [ + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1585652828216553e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.4975666999816895e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06, + 1.1213123798370361e-06 + ], + "provider_max_abs_vs_fp64": [ + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.662441253662109e-07, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.62518835067749e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 1.1026859283447266e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 9.5367431640625e-07, + 8.158385753631592e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 8.344650268554688e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.587935447692871e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 2.2687017917633057e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 1.0728836059570312e-06, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 1.0579824447631836e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 1.125037670135498e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 8.046627044677734e-07, + 5.364418029785156e-07, + 1.5720725059509277e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 8.791685104370117e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 8.642673492431641e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 6.258487701416016e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 1.0058283805847168e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 8.791685104370117e-07, + 8.344650268554688e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 9.238719940185547e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 1.1399388313293457e-06, + 5.364418029785156e-07, + 7.152557373046875e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 6.705522537231445e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 7.897615432739258e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.960464477539062e-07, + 5.364418029785156e-07, + 1.3113021850585938e-06, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07, + 5.364418029785156e-07 + ], + "rows_batch_invariant": true + }, + "perf": [ + { + "op": "timestep_mlp_fp32", + "case": "T=1", + "backend": "H3TimestepMLPCudaOp", + "bytes": 63340032, + "candidate_us": 35.472000017762184, + "candidate_gbps": 1785.634640513172, + "candidate_peak_mib": 0.05126953125, + "provider_us": 40.76800122857094, + "provider_gbps": 1553.670282849437, + "provider_peak_mib": 0.041015625 + }, + { + "op": "timestep_mlp_fp32", + "case": "T=2", + "backend": "H3TimestepMLPCudaOp", + "bytes": 63340032, + "candidate_us": 40.36800004541874, + "candidate_gbps": 1569.0653965699323, + "candidate_peak_mib": 0.1025390625, + "provider_us": 84.927998483181, + "provider_gbps": 745.808604126515, + "provider_peak_mib": 0.08203125 + }, + { + "op": "timestep_mlp_fp32", + "case": "T=3", + "backend": "H3TimestepMLPCudaOp", + "bytes": 63340032, + "candidate_us": 44.35199871659279, + "candidate_gbps": 1428.121253446544, + "candidate_peak_mib": 0.15380859375, + "provider_us": 86.94399893283844, + "provider_gbps": 728.5152831413727, + "provider_peak_mib": 0.123046875 + }, + { + "op": "timestep_mlp_fp32", + "case": "T=4", + "backend": "H3TimestepMLPCudaOp", + "bytes": 63340032, + "candidate_us": 44.943999499082565, + "candidate_gbps": 1409.310090466981, + "candidate_peak_mib": 0.205078125, + "provider_us": 88.56000006198883, + "provider_gbps": 715.2216797161726, + "provider_peak_mib": 0.1640625 + } + ] +} diff --git a/reports/experiments/h3-timestep-sinusoid-b200/figure.png b/reports/experiments/h3-timestep-sinusoid-b200/figure.png new file mode 100644 index 000000000..00810decf Binary files /dev/null and b/reports/experiments/h3-timestep-sinusoid-b200/figure.png differ diff --git a/reports/experiments/h3-timestep-sinusoid-b200/report.json b/reports/experiments/h3-timestep-sinusoid-b200/report.json new file mode 100644 index 000000000..8ab189a27 --- /dev/null +++ b/reports/experiments/h3-timestep-sinusoid-b200/report.json @@ -0,0 +1,129 @@ +{ + "kind": "h3_operator_report", + "op": "timestep_sinusoid_h3", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "reference_commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "rl_kernel_commit": "0522865203d6bfaffd9c6cd084c903f4e5864c8e", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false + }, + "accuracy": { + "contract_atol": 1e-05, + "cases": [ + { + "num_timesteps": 1, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 0.0, + "provider_max_abs_vs_fp64": 0.0 + }, + { + "num_timesteps": 2, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 5.960464477539063e-08, + "provider_max_abs_vs_fp64": 5.960464477539063e-08 + }, + { + "num_timesteps": 3, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 5.960464477539063e-08, + "provider_max_abs_vs_fp64": 5.960464477539063e-08 + }, + { + "num_timesteps": 7, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 5.960464477539063e-08, + "provider_max_abs_vs_fp64": 5.960464477539063e-08 + }, + { + "num_timesteps": 64, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 5.960464477539063e-08, + "provider_max_abs_vs_fp64": 5.960464477539063e-08 + }, + { + "num_timesteps": 1000, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 1.1920928955078125e-07, + "provider_max_abs_vs_fp64": 1.1920928955078125e-07 + }, + { + "num_timesteps": 4097, + "bitwise_equal_to_provider": true, + "max_abs_vs_fp64": 1.1920928955078125e-07, + "provider_max_abs_vs_fp64": 1.1920928955078125e-07 + } + ] + }, + "perf": [ + { + "op": "timestep_sinusoid_h3", + "case": "T=1", + "backend": "H3TimestepSinusoidCudaOp", + "bytes": 1028, + "candidate_us": 15.168000012636185, + "candidate_gbps": 0.06777426154691403, + "candidate_peak_mib": 0.0009765625, + "candidate_checked_us": 49.55200105905533, + "candidate_checked_gbps": 0.020745882669296143, + "candidate_checked_peak_mib": 0.00146484375, + "provider_us": 62.272001057863235, + "provider_gbps": 0.016508221713395416, + "provider_peak_mib": 0.0029296875 + }, + { + "op": "timestep_sinusoid_h3", + "case": "T=2", + "backend": "H3TimestepSinusoidCudaOp", + "bytes": 2056, + "candidate_us": 14.303999952971935, + "candidate_gbps": 0.1437360183696607, + "candidate_peak_mib": 0.001953125, + "candidate_checked_us": 49.47200044989586, + "candidate_checked_gbps": 0.04155886120033232, + "candidate_checked_peak_mib": 0.001953125, + "provider_us": 62.17600032687187, + "provider_gbps": 0.03306742133928188, + "provider_peak_mib": 0.00537109375 + }, + { + "op": "timestep_sinusoid_h3", + "case": "T=4", + "backend": "H3TimestepSinusoidCudaOp", + "bytes": 4112, + "candidate_us": 14.240000396966934, + "candidate_gbps": 0.2887640368939765, + "candidate_peak_mib": 0.00390625, + "candidate_checked_us": 49.18399825692177, + "candidate_checked_gbps": 0.08360442716592911, + "candidate_checked_peak_mib": 0.00390625, + "provider_us": 62.272001057863235, + "provider_gbps": 0.06603288685358166, + "provider_peak_mib": 0.01025390625 + }, + { + "op": "timestep_sinusoid_h3", + "case": "T=64", + "backend": "H3TimestepSinusoidCudaOp", + "bytes": 65792, + "candidate_us": 14.22400027513504, + "candidate_gbps": 4.625421732802616, + "candidate_peak_mib": 0.0625, + "candidate_checked_us": 49.32799935340881, + "candidate_checked_gbps": 1.3337658300032684, + "candidate_checked_peak_mib": 0.0625, + "provider_us": 70.39999961853027, + "provider_gbps": 0.9345454596093864, + "provider_peak_mib": 0.15673828125 + } + ] +} diff --git a/reports/experiments/h3-tp-adaln-3mod-b200/figure.png b/reports/experiments/h3-tp-adaln-3mod-b200/figure.png new file mode 100644 index 000000000..ed3e992fc Binary files /dev/null and b/reports/experiments/h3-tp-adaln-3mod-b200/figure.png differ diff --git a/reports/experiments/h3-tp-adaln-3mod-b200/report.json b/reports/experiments/h3-tp-adaln-3mod-b200/report.json new file mode 100644 index 000000000..28090765f --- /dev/null +++ b/reports/experiments/h3-tp-adaln-3mod-b200/report.json @@ -0,0 +1,1771 @@ +{ + "kind": "h3_ws2_report", + "op": "tp_adaln_3mod", + "rfc": "RL-Align/RL-Kernel#420", + "model_revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "rl_kernel_commit": "49479963212460d542a694d3fbd79011c766b643", + "tracked_tree_dirty": false, + "environment": { + "gpu": "NVIDIA B200", + "capability": [ + 10, + 0 + ], + "torch": "2.13.0+cu130", + "cuda": "13.0", + "python": "3.12.14", + "tf32": false, + "gpus": 8 + }, + "weights": [ + "transformer_blocks.0.adaln_proj.linear.weight", + "transformer_blocks.0.adaln_proj.linear.bias" + ], + "runs": [ + { + "tp": 1, + "ranks": [ + { + "rank": 0, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 103.53599861264229, + "tp_forward_backward_us": 1937.5039935112, + "ws1_forward_us": 100.27200356125832, + "ws1_forward_backward_us": 1904.703974723816 + }, + { + "num_timesteps": 2, + "tp_forward_us": 103.88800129294395, + "tp_forward_backward_us": 2020.3039646148682, + "ws1_forward_us": 100.38399696350098, + "ws1_forward_backward_us": 1987.5360131263733 + }, + { + "num_timesteps": 3, + "tp_forward_us": 103.2319962978363, + "tp_forward_backward_us": 2281.4719676971436, + "ws1_forward_us": 100.12800246477127, + "ws1_forward_backward_us": 2246.367931365967 + }, + { + "num_timesteps": 4, + "tp_forward_us": 103.5039983689785, + "tp_forward_backward_us": 2131.103992462158, + "ws1_forward_us": 100.44799745082855, + "ws1_forward_backward_us": 2093.3759212493896 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 1, + "rank": 0, + "columns": [ + 0, + 96768 + ], + "slots": [ + [ + 0, + 0, + 0, + 5376 + ], + [ + 0, + 1, + 0, + 5376 + ], + [ + 0, + 2, + 0, + 5376 + ], + [ + 0, + 3, + 0, + 5376 + ], + [ + 0, + 4, + 0, + 5376 + ], + [ + 0, + 5, + 0, + 5376 + ], + [ + 1, + 0, + 0, + 5376 + ], + [ + 1, + 1, + 0, + 5376 + ], + [ + 1, + 2, + 0, + 5376 + ], + [ + 1, + 3, + 0, + 5376 + ], + [ + 1, + 4, + 0, + 5376 + ], + [ + 1, + 5, + 0, + 5376 + ], + [ + 2, + 0, + 0, + 5376 + ], + [ + 2, + 1, + 0, + 5376 + ], + [ + 2, + 2, + 0, + 5376 + ], + [ + 2, + 3, + 0, + 5376 + ], + [ + 2, + 4, + 0, + 5376 + ], + [ + 2, + 5, + 0, + 5376 + ] + ], + "fallback": null + } + } + ] + }, + { + "tp": 2, + "ranks": [ + { + "rank": 0, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 120.78399956226349, + "tp_forward_backward_us": 1487.9680275917053, + "ws1_forward_us": 100.73599964380264, + "ws1_forward_backward_us": 1912.6240015029907 + }, + { + "num_timesteps": 2, + "tp_forward_us": 113.96800354123116, + "tp_forward_backward_us": 1553.8560152053833, + "ws1_forward_us": 101.45599767565727, + "ws1_forward_backward_us": 1974.3199944496155 + }, + { + "num_timesteps": 3, + "tp_forward_us": 109.58399996161461, + "tp_forward_backward_us": 1768.7840461730957, + "ws1_forward_us": 101.3919971883297, + "ws1_forward_backward_us": 2222.848057746887 + }, + { + "num_timesteps": 4, + "tp_forward_us": 111.85600236058235, + "tp_forward_backward_us": 1820.3840255737305, + "ws1_forward_us": 101.21600329875946, + "ws1_forward_backward_us": 2071.5359449386597 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 2, + "rank": 0, + "columns": [ + 0, + 48384 + ], + "slots": [ + [ + 0, + 0, + 0, + 5376 + ], + [ + 0, + 1, + 0, + 5376 + ], + [ + 0, + 2, + 0, + 5376 + ], + [ + 0, + 3, + 0, + 5376 + ], + [ + 0, + 4, + 0, + 5376 + ], + [ + 0, + 5, + 0, + 5376 + ], + [ + 1, + 0, + 0, + 5376 + ], + [ + 1, + 1, + 0, + 5376 + ], + [ + 1, + 2, + 0, + 5376 + ] + ], + "fallback": null + } + }, + { + "rank": 1, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 120.89599668979645, + "tp_forward_backward_us": 1489.9200201034546, + "ws1_forward_us": 100.28800368309021, + "ws1_forward_backward_us": 1906.2560200691223 + }, + { + "num_timesteps": 2, + "tp_forward_us": 114.88000303506851, + "tp_forward_backward_us": 1559.6320033073425, + "ws1_forward_us": 100.00000149011612, + "ws1_forward_backward_us": 1850.7359623908997 + }, + { + "num_timesteps": 3, + "tp_forward_us": 109.37599837779999, + "tp_forward_backward_us": 1771.2000012397766, + "ws1_forward_us": 100.09600222110748, + "ws1_forward_backward_us": 2145.6159353256226 + }, + { + "num_timesteps": 4, + "tp_forward_us": 112.0000034570694, + "tp_forward_backward_us": 1826.9280195236206, + "ws1_forward_us": 100.08000209927559, + "ws1_forward_backward_us": 1989.3759489059448 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 2, + "rank": 1, + "columns": [ + 48384, + 96768 + ], + "slots": [ + [ + 1, + 3, + 0, + 5376 + ], + [ + 1, + 4, + 0, + 5376 + ], + [ + 1, + 5, + 0, + 5376 + ], + [ + 2, + 0, + 0, + 5376 + ], + [ + 2, + 1, + 0, + 5376 + ], + [ + 2, + 2, + 0, + 5376 + ], + [ + 2, + 3, + 0, + 5376 + ], + [ + 2, + 4, + 0, + 5376 + ], + [ + 2, + 5, + 0, + 5376 + ] + ], + "fallback": null + } + } + ] + }, + { + "tp": 4, + "ranks": [ + { + "rank": 0, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 113.42399939894676, + "tp_forward_backward_us": 1185.85604429245, + "ws1_forward_us": 100.22400319576263, + "ws1_forward_backward_us": 1903.4239649772644 + }, + { + "num_timesteps": 2, + "tp_forward_us": 106.46399855613708, + "tp_forward_backward_us": 1289.5520329475403, + "ws1_forward_us": 100.14400258660316, + "ws1_forward_backward_us": 1922.7680563926697 + }, + { + "num_timesteps": 3, + "tp_forward_us": 110.47999933362007, + "tp_forward_backward_us": 1472.544014453888, + "ws1_forward_us": 100.832000374794, + "ws1_forward_backward_us": 2194.5440769195557 + }, + { + "num_timesteps": 4, + "tp_forward_us": 115.15199765563011, + "tp_forward_backward_us": 1565.2480125427246, + "ws1_forward_us": 100.43199732899666, + "ws1_forward_backward_us": 2023.8720178604126 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 4, + "rank": 0, + "columns": [ + 0, + 24192 + ], + "slots": [ + [ + 0, + 0, + 0, + 5376 + ], + [ + 0, + 1, + 0, + 5376 + ], + [ + 0, + 2, + 0, + 5376 + ], + [ + 0, + 3, + 0, + 5376 + ], + [ + 0, + 4, + 0, + 2688 + ] + ], + "fallback": null + } + }, + { + "rank": 1, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 112.84800246357918, + "tp_forward_backward_us": 1189.1520023345947, + "ws1_forward_us": 100.0640019774437, + "ws1_forward_backward_us": 1912.6240015029907 + }, + { + "num_timesteps": 2, + "tp_forward_us": 106.22400045394897, + "tp_forward_backward_us": 1291.3440465927124, + "ws1_forward_us": 99.82400014996529, + "ws1_forward_backward_us": 1985.8239889144897 + }, + { + "num_timesteps": 3, + "tp_forward_us": 110.22399738430977, + "tp_forward_backward_us": 1476.9439697265625, + "ws1_forward_us": 100.20800307393074, + "ws1_forward_backward_us": 2261.5840435028076 + }, + { + "num_timesteps": 4, + "tp_forward_us": 114.81600254774094, + "tp_forward_backward_us": 1567.0239925384521, + "ws1_forward_us": 100.03200173377991, + "ws1_forward_backward_us": 2092.7200317382812 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 4, + "rank": 1, + "columns": [ + 24192, + 48384 + ], + "slots": [ + [ + 0, + 4, + 2688, + 5376 + ], + [ + 0, + 5, + 0, + 5376 + ], + [ + 1, + 0, + 0, + 5376 + ], + [ + 1, + 1, + 0, + 5376 + ], + [ + 1, + 2, + 0, + 5376 + ] + ], + "fallback": null + } + }, + { + "rank": 2, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 113.00799995660782, + "tp_forward_backward_us": 1192.1120285987854, + "ws1_forward_us": 100.44799745082855, + "ws1_forward_backward_us": 1802.9919862747192 + }, + { + "num_timesteps": 2, + "tp_forward_us": 106.47999867796898, + "tp_forward_backward_us": 1288.144052028656, + "ws1_forward_us": 100.46399757266045, + "ws1_forward_backward_us": 1897.9679942131042 + }, + { + "num_timesteps": 3, + "tp_forward_us": 110.3999987244606, + "tp_forward_backward_us": 1477.679967880249, + "ws1_forward_us": 100.47999769449234, + "ws1_forward_backward_us": 2241.7759895324707 + }, + { + "num_timesteps": 4, + "tp_forward_us": 115.10399729013443, + "tp_forward_backward_us": 1566.096007823944, + "ws1_forward_us": 100.70399940013885, + "ws1_forward_backward_us": 2089.9680852890015 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 4, + "rank": 2, + "columns": [ + 48384, + 72576 + ], + "slots": [ + [ + 1, + 3, + 0, + 5376 + ], + [ + 1, + 4, + 0, + 5376 + ], + [ + 1, + 5, + 0, + 5376 + ], + [ + 2, + 0, + 0, + 5376 + ], + [ + 2, + 1, + 0, + 2688 + ] + ], + "fallback": null + } + }, + { + "rank": 3, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 113.52000012993813, + "tp_forward_backward_us": 1189.184010028839, + "ws1_forward_us": 100.63999891281128, + "ws1_forward_backward_us": 1900.8640050888062 + }, + { + "num_timesteps": 2, + "tp_forward_us": 106.6720001399517, + "tp_forward_backward_us": 1289.3120050430298, + "ws1_forward_us": 100.54399818181992, + "ws1_forward_backward_us": 1980.672001838684 + }, + { + "num_timesteps": 3, + "tp_forward_us": 110.57600006461143, + "tp_forward_backward_us": 1475.9039878845215, + "ws1_forward_us": 100.38399696350098, + "ws1_forward_backward_us": 2248.5759258270264 + }, + { + "num_timesteps": 4, + "tp_forward_us": 115.23199826478958, + "tp_forward_backward_us": 1566.2879943847656, + "ws1_forward_us": 100.28800368309021, + "ws1_forward_backward_us": 2051.9360303878784 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 4, + "rank": 3, + "columns": [ + 72576, + 96768 + ], + "slots": [ + [ + 2, + 1, + 2688, + 5376 + ], + [ + 2, + 2, + 0, + 5376 + ], + [ + 2, + 3, + 0, + 5376 + ], + [ + 2, + 4, + 0, + 5376 + ], + [ + 2, + 5, + 0, + 5376 + ] + ], + "fallback": null + } + } + ] + }, + { + "tp": 8, + "ranks": [ + { + "rank": 0, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 122.92800098657608, + "tp_forward_backward_us": 1037.1840000152588, + "ws1_forward_us": 100.65599903464317, + "ws1_forward_backward_us": 1905.2960276603699 + }, + { + "num_timesteps": 2, + "tp_forward_us": 115.31199887394905, + "tp_forward_backward_us": 1155.1520228385925, + "ws1_forward_us": 100.76799988746643, + "ws1_forward_backward_us": 1995.8879947662354 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.91199681162834, + "tp_forward_backward_us": 1318.4799551963806, + "ws1_forward_us": 100.3040000796318, + "ws1_forward_backward_us": 2250.5760192871094 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.98399752378464, + "tp_forward_backward_us": 1444.7839856147766, + "ws1_forward_us": 118.60800161957741, + "ws1_forward_backward_us": 2055.567979812622 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 0, + "columns": [ + 0, + 12096 + ], + "slots": [ + [ + 0, + 0, + 0, + 5376 + ], + [ + 0, + 1, + 0, + 5376 + ], + [ + 0, + 2, + 0, + 1344 + ] + ], + "fallback": null + } + }, + { + "rank": 1, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 122.81600013375282, + "tp_forward_backward_us": 1038.0959510803223, + "ws1_forward_us": 100.27200356125832, + "ws1_forward_backward_us": 1905.4880142211914 + }, + { + "num_timesteps": 2, + "tp_forward_us": 115.00799655914307, + "tp_forward_backward_us": 1149.0240097045898, + "ws1_forward_us": 100.17600283026695, + "ws1_forward_backward_us": 2000.256061553955 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.88000029325485, + "tp_forward_backward_us": 1321.0880160331726, + "ws1_forward_us": 100.41599720716476, + "ws1_forward_backward_us": 2252.2239685058594 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.26399949193001, + "tp_forward_backward_us": 1448.0960369110107, + "ws1_forward_us": 118.51200088858604, + "ws1_forward_backward_us": 2069.551944732666 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 1, + "columns": [ + 12096, + 24192 + ], + "slots": [ + [ + 0, + 2, + 1344, + 5376 + ], + [ + 0, + 3, + 0, + 5376 + ], + [ + 0, + 4, + 0, + 2688 + ] + ], + "fallback": null + } + }, + { + "rank": 2, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 123.10399860143661, + "tp_forward_backward_us": 1037.7439856529236, + "ws1_forward_us": 100.16000270843506, + "ws1_forward_backward_us": 1916.0799980163574 + }, + { + "num_timesteps": 2, + "tp_forward_us": 114.8960031569004, + "tp_forward_backward_us": 1151.199996471405, + "ws1_forward_us": 100.36799684166908, + "ws1_forward_backward_us": 2007.1839094161987 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.95999717712402, + "tp_forward_backward_us": 1318.4320330619812, + "ws1_forward_us": 100.76799988746643, + "ws1_forward_backward_us": 2272.8641033172607 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.40800058841705, + "tp_forward_backward_us": 1448.0479955673218, + "ws1_forward_us": 100.8480004966259, + "ws1_forward_backward_us": 2094.8160886764526 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 2, + "columns": [ + 24192, + 36288 + ], + "slots": [ + [ + 0, + 4, + 2688, + 5376 + ], + [ + 0, + 5, + 0, + 5376 + ], + [ + 1, + 0, + 0, + 4032 + ] + ], + "fallback": null + } + }, + { + "rank": 3, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 123.23199957609177, + "tp_forward_backward_us": 1038.6560559272766, + "ws1_forward_us": 100.36799684166908, + "ws1_forward_backward_us": 1913.103997707367 + }, + { + "num_timesteps": 2, + "tp_forward_us": 115.34399911761284, + "tp_forward_backward_us": 1151.1359810829163, + "ws1_forward_us": 100.3199964761734, + "ws1_forward_backward_us": 1992.512047290802 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.79999968409538, + "tp_forward_backward_us": 1319.3119764328003, + "ws1_forward_us": 100.44799745082855, + "ws1_forward_backward_us": 2261.47198677063 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.99999764561653, + "tp_forward_backward_us": 1447.8240013122559, + "ws1_forward_us": 100.68799927830696, + "ws1_forward_backward_us": 2095.2640771865845 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 3, + "columns": [ + 36288, + 48384 + ], + "slots": [ + [ + 1, + 0, + 4032, + 5376 + ], + [ + 1, + 1, + 0, + 5376 + ], + [ + 1, + 2, + 0, + 5376 + ] + ], + "fallback": null + } + }, + { + "rank": 4, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 122.99199774861336, + "tp_forward_backward_us": 1036.7679595947266, + "ws1_forward_us": 101.42399743199348, + "ws1_forward_backward_us": 1913.152039051056 + }, + { + "num_timesteps": 2, + "tp_forward_us": 115.77600240707397, + "tp_forward_backward_us": 1149.4719982147217, + "ws1_forward_us": 101.50399804115295, + "ws1_forward_backward_us": 1987.5360131263733 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.64000219106674, + "tp_forward_backward_us": 1319.6800351142883, + "ws1_forward_us": 101.43999755382538, + "ws1_forward_backward_us": 2259.4879865646362 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.95199728012085, + "tp_forward_backward_us": 1449.1999745368958, + "ws1_forward_us": 101.53599828481674, + "ws1_forward_backward_us": 2089.1839265823364 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 4, + "columns": [ + 48384, + 60480 + ], + "slots": [ + [ + 1, + 3, + 0, + 5376 + ], + [ + 1, + 4, + 0, + 5376 + ], + [ + 1, + 5, + 0, + 1344 + ] + ], + "fallback": null + } + }, + { + "rank": 5, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 122.22399935126305, + "tp_forward_backward_us": 1038.1120443344116, + "ws1_forward_us": 100.92800110578537, + "ws1_forward_backward_us": 1924.8319864273071 + }, + { + "num_timesteps": 2, + "tp_forward_us": 114.31999877095222, + "tp_forward_backward_us": 1152.0960330963135, + "ws1_forward_us": 101.24800354242325, + "ws1_forward_backward_us": 1999.72802400589 + }, + { + "num_timesteps": 3, + "tp_forward_us": 119.90399658679962, + "tp_forward_backward_us": 1320.2400207519531, + "ws1_forward_us": 101.6319990158081, + "ws1_forward_backward_us": 2250.3679990768433 + }, + { + "num_timesteps": 4, + "tp_forward_us": 120.86400017142296, + "tp_forward_backward_us": 1448.0000138282776, + "ws1_forward_us": 101.31199657917023, + "ws1_forward_backward_us": 2092.8800106048584 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 5, + "columns": [ + 60480, + 72576 + ], + "slots": [ + [ + 1, + 5, + 1344, + 5376 + ], + [ + 2, + 0, + 0, + 5376 + ], + [ + 2, + 1, + 0, + 2688 + ] + ], + "fallback": null + } + }, + { + "rank": 6, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 121.95199728012085, + "tp_forward_backward_us": 1039.8079752922058, + "ws1_forward_us": 101.67999938130379, + "ws1_forward_backward_us": 1919.647991657257 + }, + { + "num_timesteps": 2, + "tp_forward_us": 114.17599767446518, + "tp_forward_backward_us": 1149.5360136032104, + "ws1_forward_us": 101.80800035595894, + "ws1_forward_backward_us": 2007.6640844345093 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.11199817061424, + "tp_forward_backward_us": 1320.1760053634644, + "ws1_forward_us": 101.96800157427788, + "ws1_forward_backward_us": 2253.376007080078 + }, + { + "num_timesteps": 4, + "tp_forward_us": 121.74400314688683, + "tp_forward_backward_us": 1447.8400349617004, + "ws1_forward_us": 101.87200084328651, + "ws1_forward_backward_us": 2092.3839807510376 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 6, + "columns": [ + 72576, + 84672 + ], + "slots": [ + [ + 2, + 1, + 2688, + 5376 + ], + [ + 2, + 2, + 0, + 5376 + ], + [ + 2, + 3, + 0, + 4032 + ] + ], + "fallback": null + } + }, + { + "rank": 7, + "equality": [ + { + "num_timesteps": 1, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 2, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 3, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + }, + { + "num_timesteps": 4, + "table": true, + "d_temb": true, + "d_weight_shard": true, + "d_bias_shard": true + } + ], + "perf": [ + { + "num_timesteps": 1, + "tp_forward_us": 121.85600027441978, + "tp_forward_backward_us": 1035.5199575424194, + "ws1_forward_us": 101.29599645733833, + "ws1_forward_backward_us": 1919.215977191925 + }, + { + "num_timesteps": 2, + "tp_forward_us": 114.54400047659874, + "tp_forward_backward_us": 1148.800015449524, + "ws1_forward_us": 101.53599828481674, + "ws1_forward_backward_us": 1993.5839772224426 + }, + { + "num_timesteps": 3, + "tp_forward_us": 120.06399780511856, + "tp_forward_backward_us": 1317.9360032081604, + "ws1_forward_us": 101.31199657917023, + "ws1_forward_backward_us": 2245.6640005111694 + }, + { + "num_timesteps": 4, + "tp_forward_us": 120.67200243473053, + "tp_forward_backward_us": 1448.5120177268982, + "ws1_forward_us": 101.24800354242325, + "ws1_forward_backward_us": 2085.088014602661 + } + ], + "readback": { + "op": "tp_adaln_3mod", + "kernel_id": "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]", + "backward_kernel_id": "rl_engine._C.h3_det_linear_backward_weight+rl_engine._C.h3_det_linear_backward_input_partials+all_gather+rl_engine._C.h3_det_linear_fold_chunks", + "contract": "h3-det-linear-v1", + "collective_backend": "cuda_ipc_fixed_tree", + "collective_ops": [ + "all_gather" + ], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": 8, + "rank": 7, + "columns": [ + 84672, + 96768 + ], + "slots": [ + [ + 2, + 3, + 4032, + 5376 + ], + [ + 2, + 4, + 0, + 5376 + ], + [ + 2, + 5, + 0, + 5376 + ] + ], + "fallback": null + } + } + ] + } + ] +} diff --git a/rl_engine/_C.pyi b/rl_engine/_C.pyi index fb3c4379c..40fb56213 100644 --- a/rl_engine/_C.pyi +++ b/rl_engine/_C.pyi @@ -325,3 +325,153 @@ def deterministic_collective_rocm_ipc_all_gather_input( input: torch.Tensor, output: torch.Tensor, ) -> None: ... +def h3_timestep_sinusoid_forward( + timestep: torch.Tensor, + num_channels: int = 256, + max_period: float = 10000.0, + check_range: bool = True, +) -> torch.Tensor: + """Return FP32 CUDA cosine-then-sine features of shape ``(T, num_channels)``. + + Require a nonempty 1-D FP32 CUDA timestep tensor, positive even channels + and positive ``max_period``. Check finite values in ``[0, 1]`` when enabled; + the Python operator supplies the analytic backward for this native forward. + """ + ... + +def h3_det_linear_forward( + x: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor | None = None, + activation: int = 0, + save_pre_activation: bool = False, +) -> list[torch.Tensor]: + """Return a deterministic CUDA projection and optional FP32 pre-activation. + + Require same-device/dtype FP32 or BF16 contiguous, 16-byte-aligned matrices + ``x (T, K)`` and ``weight (N, K)``, ``T > 0`` and ``K`` divisible by the + 16-byte vector width. Optional contiguous bias is ``(N,)`` in that dtype. + Activation 0 is identity and 1 is SiLU; BF16 forward requires SM80 or newer. + Return ``[out]`` in the input dtype or ``[out, pre]`` when saving activation. + """ + ... + +def h3_det_linear_backward_input( + grad: torch.Tensor, weight: torch.Tensor, out_dtype: torch.dtype +) -> torch.Tensor: + """Return a deterministic ``(T, K)`` CUDA input gradient in ``out_dtype``. + + Require contiguous FP32 ``grad (T, N)`` with ``T > 0`` and same-device + contiguous, 16-byte-aligned FP32 or BF16 ``weight (N, K)``. Output dtype + must be FP32 or BF16; output-column chunks accumulate in fixed order. + """ + ... + +def h3_det_linear_backward_input_partials( + grad: torch.Tensor, weight: torch.Tensor +) -> torch.Tensor: ... +def h3_det_linear_fold_chunks(partial: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: ... +def h3_det_linear_backward_weight( + grad: torch.Tensor, x: torch.Tensor, w_dtype: torch.dtype, with_bias: bool = True +) -> list[torch.Tensor]: + """Return CUDA weight and optional bias gradients using an ascending row fold. + + Require contiguous FP32 ``grad (T, N)`` with ``T > 0`` and same-device + contiguous, 16-byte-aligned FP32 or BF16 ``x (T, K)``. Return ``[dweight]`` + with shape ``(N, K)``, or ``[dweight, dbias]`` with ``dbias (N,)`` when + requested, in FP32 or BF16 as specified by ``w_dtype``. + """ + ... + +def h3_adaln_row_gather_forward( + rows: torch.Tensor, + timestep_indices: torch.Tensor, + token_tags: torch.Tensor, + chunks: int = 6, + modality_num: int = 3, +) -> torch.Tensor: ... +def h3_adaln_row_gather_backward( + grad: torch.Tensor, + sorted_pos: torch.Tensor, + tile_begin: torch.Tensor, + tile_end: torch.Tensor, + seg_first_tile: torch.Tensor, + out_dtype: torch.dtype, +) -> torch.Tensor: ... +def h3_rmsnorm_forward( + x: torch.Tensor, + weight: torch.Tensor, + eps: float, + shift: torch.Tensor | None = None, + scale: torch.Tensor | None = None, + index: torch.Tensor | None = None, +) -> list[torch.Tensor]: ... +def h3_rmsnorm_backward( + grad: torch.Tensor, + x: torch.Tensor, + weight: torch.Tensor, + rstd: torch.Tensor, + shift: torch.Tensor | None = None, + scale: torch.Tensor | None = None, + index: torch.Tensor | None = None, + sorted_pos: torch.Tensor | None = None, + tile_begin: torch.Tensor | None = None, + tile_end: torch.Tensor | None = None, + seg_first_tile: torch.Tensor | None = None, +) -> list[torch.Tensor]: ... +def h3_gate_residual_forward( + residual: torch.Tensor, y: torch.Tensor, gate: torch.Tensor, index: torch.Tensor +) -> torch.Tensor: ... +def h3_gate_residual_backward( + 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: torch.Tensor, +) -> list[torch.Tensor]: ... +def h3_rmsnorm_backward_partials( + grad: torch.Tensor, + x: torch.Tensor, + weight: torch.Tensor, + rstd: torch.Tensor, + shift: torch.Tensor | None, + scale: torch.Tensor | None, + index: torch.Tensor | None, + dw_rows: torch.Tensor, + dw_begin: torch.Tensor, + dw_end: torch.Tensor, + seg_rows: torch.Tensor | None = None, + seg_begin: torch.Tensor | None = None, + seg_end: torch.Tensor | None = None, +) -> list[torch.Tensor]: ... +def h3_rmsnorm_fold_partials( + dw_partial: torch.Tensor, + weight: torch.Tensor, + seg_partial: torch.Tensor | None = None, + seg_first_tile: torch.Tensor | None = None, +) -> list[torch.Tensor]: ... +def h3_gate_grad_partials( + grad: torch.Tensor, + y: torch.Tensor, + rows: torch.Tensor, + tile_begin: torch.Tensor, + tile_end: torch.Tensor, +) -> torch.Tensor: ... +def h3_gate_grad_fold( + partial: torch.Tensor, seg_first_tile: torch.Tensor, dtype: torch.dtype +) -> torch.Tensor: ... +def h3_rmsnorm_backward_dx( + grad: torch.Tensor, + x: torch.Tensor, + weight: torch.Tensor, + rstd: torch.Tensor, + shift: torch.Tensor | None = None, + scale: torch.Tensor | None = None, + index: torch.Tensor | None = None, +) -> torch.Tensor: ... +def h3_gate_residual_backward_dy( + grad: torch.Tensor, gate: torch.Tensor, index: torch.Tensor +) -> torch.Tensor: ... diff --git a/rl_engine/backends/cuda/model_specific/__init__.py b/rl_engine/backends/cuda/model_specific/__init__.py new file mode 100644 index 000000000..3b09d3c11 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CUDA implementations with model-specific arithmetic contracts.""" diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/__init__.py b/rl_engine/backends/cuda/model_specific/minimax_h3/__init__.py new file mode 100644 index 000000000..86cf4c9da --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_modulation.py b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_modulation.py new file mode 100644 index 000000000..763dba546 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_modulation.py @@ -0,0 +1,133 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Fused H3 AdaLN modulation: ``adaln_projection_3mod`` followed by ``adaln_row_gather``. + +The forward is exactly the two RFC #420 ops back to back (same kernels, same +bits). The difference is the backward. With two separate ops, autograd hands +the gather's table gradient to the projection in the table's dtype (BF16), +which rounds a long FP32 segment sum once more before it reaches the +projection weights and ``temb``. Here the segment sum stays FP32 into the +projection backward, so the parameter gradients carry no extra BF16 rounding. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C +from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + _segment_tiles, + adaln_row_gather_available, +) +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + det_linear_available, + det_linear_backward_input, + det_linear_backward_weight, + det_linear_forward, + silu_backward_fp32, +) +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM +from rl_engine.reference.minimax_h3.adaln_projection import validate_h3_adaln_projection +from rl_engine.reference.minimax_h3.adaln_row_gather import ( + h3_adaln_indices, + validate_h3_adaln_row_gather, +) + +KERNEL_ID = ( + "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward" + "+rl_engine._C.h3_adaln_row_gather_forward" +) +BACKWARD_IMPL = "fp32_segment_sum_into_h3_det_linear_backward" +BACKWARD_KERNEL_ID = ( + "rl_engine._C.h3_adaln_row_gather_backward[fp32]" + "+rl_engine._C.h3_det_linear_backward_weight" + "+rl_engine._C.h3_det_linear_backward_input" + "+rl_engine.backends.cuda.model_specific.minimax_h3.det_linear.silu_backward_fp32" +) + + +class _H3AdaLNModulationCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, temb, weight, bias, timestep_indices, token_tags, hidden): + temb = temb.contiguous() + act = F.silu(temb).to(weight.dtype) + (table,) = det_linear_forward(act, weight, bias) + rows = table.view(-1, H3_ADALN_CHUNKS * hidden) + packed = _C.h3_adaln_row_gather_forward( + rows, timestep_indices, token_tags, H3_ADALN_CHUNKS, H3_MODALITY_NUM + ) + ctx.save_for_backward(temb, act, weight, timestep_indices, token_tags) + ctx.table_shape = table.shape + return packed + + @staticmethod + def backward(ctx, grad_packed): + temb, act, weight, timestep_indices, token_tags = ctx.saved_tensors + need_temb, need_weight, need_bias = ctx.needs_input_grad[:3] + num_timesteps, width = ctx.table_shape + index = h3_adaln_indices(timestep_indices, token_tags) + tiles = _segment_tiles(index, num_timesteps * H3_MODALITY_NUM) + d_rows = _C.h3_adaln_row_gather_backward(grad_packed.contiguous(), *tiles, torch.float32) + d_table = d_rows.view(num_timesteps, width) + d_temb = dw = db = None + if need_weight or need_bias: + dw, db = det_linear_backward_weight(d_table, act, weight.dtype) + if need_temb: + d_temb = silu_backward_fp32( + det_linear_backward_input(d_table, weight, torch.float32), temb + ) + record_backward( + "adaln_modulation_3mod", + kernel_id=BACKWARD_KERNEL_ID, + impl=BACKWARD_IMPL, + family="cuda", + ) + return ( + d_temb, + dw if need_weight else None, + db if need_bias else None, + None, + None, + None, + ) + + +class H3AdaLNModulationCudaOp: + """``adaln_projection_3mod`` + ``adaln_row_gather`` with an FP32 table gradient. + + Returns the six ``(S, H)`` modulation tensors for every packed position, + bitwise equal to running the two ops separately. + """ + + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + if not (det_linear_available() and adaln_row_gather_available()): + raise RuntimeError("rl_engine._C lacks the h3_det_linear_* / h3_adaln_row_gather_* ops") + + def __call__(self, temb, weight, bias, timestep_indices, token_tags, *, check_range=True): + hidden = validate_h3_adaln_projection(temb, weight, bias) + if not temb.is_cuda: + raise ValueError("H3AdaLNModulationCudaOp needs CUDA tensors") + _validate_indices(timestep_indices, token_tags, temb.shape[0], check_range) + packed = _H3AdaLNModulationCuda.apply( + temb, + weight, + bias, + timestep_indices.contiguous(), + token_tags.contiguous(), + hidden, + ) + return tuple(packed.unbind(0)) + + +def _validate_indices(timestep_indices, token_tags, num_timesteps, check_range): + # Reuse the gather's checks against a stand-in of the right row count. + stand_in = torch.empty( + (num_timesteps * H3_MODALITY_NUM, H3_ADALN_CHUNKS), device=timestep_indices.device + ) + validate_h3_adaln_row_gather(stand_in, timestep_indices, token_tags, check_range=check_range) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_projection.py b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_projection.py new file mode 100644 index 000000000..6c9f5bf49 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_projection.py @@ -0,0 +1,132 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 three-modality AdaLN projection (RFC #420 ``adaln_projection_3mod``).""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + ACT_NONE, + CONTRACT, + det_linear_available, + det_linear_backward_input, + det_linear_backward_weight, + det_linear_forward, + silu_backward_fp32, +) +from rl_engine.reference.minimax_h3.adaln_projection import ( + split_adaln_table, + validate_h3_adaln_projection, +) + +KERNEL_ID = "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward" +BACKWARD_IMPL = "h3_det_linear_v1_chunked_dinput_ascending_row_dweight" +BACKWARD_KERNEL_ID = ( + "rl_engine._C.h3_det_linear_backward_weight" + "+rl_engine._C.h3_det_linear_backward_input" + "+rl_engine.backends.cuda.model_specific.minimax_h3.det_linear.silu_backward_fp32" +) + + +class _H3AdaLNProjectionCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, temb, weight, bias): + """Project FP32 ``(T, D)`` embeddings to a weight-dtype ``(T, 18H)`` table. + + Save the embedding, rounded SiLU activation and weight for the VJP. + """ + + temb = temb.contiguous() + # Declared mixed-precision boundary: SiLU at temb's FP32 precision, + # then exactly one rounding to the projection dtype. These are the + # provider's own elementwise ops, so the activation is bitwise equal. + act = F.silu(temb).to(weight.dtype) + (table,) = det_linear_forward(act, weight, bias, activation=ACT_NONE) + ctx.save_for_backward(temb, act, weight) + return table + + @staticmethod + def backward(ctx, grad_table): + """Return requested embedding, weight and bias gradients in their input dtypes. + + Accumulate projection gradients deterministically and evaluate the + straight-through cast and SiLU VJP in FP32 for the embedding gradient. + """ + + temb, act, weight = ctx.saved_tensors + need_temb, need_weight, need_bias = ctx.needs_input_grad + grad_table = grad_table.float().contiguous() + d_temb = dw = db = None + if need_weight or need_bias: + dw, db = det_linear_backward_weight(grad_table, act, weight.dtype) + if need_temb: + # The cast's VJP is the identity; the SiLU VJP runs in FP32. + d_act = det_linear_backward_input(grad_table, weight, torch.float32) + d_temb = silu_backward_fp32(d_act, temb) + record_backward( + "adaln_projection_3mod", + kernel_id=BACKWARD_KERNEL_ID, + impl=BACKWARD_IMPL, + family="cuda", + ) + return d_temb, dw if need_weight else None, db if need_bias else None + + +class H3AdaLNProjectionCudaOp: + """CUDA candidate: FP32 SiLU, one BF16 cast, deterministic GEMV, six views. + + The projection weight (96768 x 2688 BF16, 520 MB per block) is streamed + once per call; every output element follows contract ``h3-det-linear-v1``, + so a timestep's 3 x 6 modulation rows do not depend on the other + timesteps in the call. The six outputs are views of one ``(T, 6H*3)`` + table, laid out exactly like diffusers' ``view(-1, 6H).chunk(6)``. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + contract = CONTRACT + + def __init__(self) -> None: + """Require all deterministic linear symbols in the CUDA extension.""" + + if not det_linear_available(): + raise RuntimeError( + "rl_engine._C lacks h3_det_linear_*; rebuild with csrc/cuda/h3/det_linear.cu" + ) + + def __call__(self, temb, weight, bias): + """Return six weight-dtype ``(3T, H)`` modulation views for CUDA inputs.""" + + return self.forward(temb, weight, bias) + + def forward(self, temb, weight, bias) -> tuple[torch.Tensor, ...]: + """Return six ``(3T, H)`` views after FP32 SiLU and a deterministic projection. + + Require CUDA FP32 ``temb`` of shape ``(T, D)`` with ``T > 0`` and + same-device FP32 or BF16 weight/bias of shapes ``(18H, D)``/``(18H,)``. + Outputs share one table in the weight dtype; autograd covers all inputs. + """ + + hidden = validate_h3_adaln_projection(temb, weight, bias) + if not temb.is_cuda: + raise ValueError("H3AdaLNProjectionCudaOp needs CUDA tensors") + table = _H3AdaLNProjectionCuda.apply(temb, weight, bias) + return split_adaln_table(table, hidden) + + def forward_table(self, temb, weight, bias) -> torch.Tensor: + """The raw ``(T, 6H*3)`` projection, for consumers that gather it directly.""" + + validate_h3_adaln_projection(temb, weight, bias) + if not temb.is_cuda: + raise ValueError("H3AdaLNProjectionCudaOp needs CUDA tensors") + return _H3AdaLNProjectionCuda.apply(temb, weight, bias) + + def forward_fp32(self, temb, weight, bias): + """Run the CUDA projection with its declared weight-dtype output boundary.""" + + return self.forward(temb, weight, bias) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_row_gather.py b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_row_gather.py new file mode 100644 index 000000000..92835f0cc --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/adaln_row_gather.py @@ -0,0 +1,113 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 AdaLN row gather (RFC #420 ``adaln_row_gather``).""" + +from __future__ import annotations + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C, _EXT_AVAILABLE +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM +from rl_engine.reference.minimax_h3.adaln_row_gather import ( + h3_adaln_indices, + validate_h3_adaln_row_gather, +) + +KERNEL_ID = "rl_engine._C.h3_adaln_row_gather_forward" +BACKWARD_IMPL = "stable_sorted_segment_tiles_fp32_ordered_fold" +BACKWARD_KERNEL_ID = "torch.argsort[stable]+rl_engine._C.h3_adaln_row_gather_backward" +# Positions per backward tile; part of the fixed summation order. +BACKWARD_TILE = 256 + + +def adaln_row_gather_available() -> bool: + return bool( + _EXT_AVAILABLE + and hasattr(_C, "h3_adaln_row_gather_forward") + and hasattr(_C, "h3_adaln_row_gather_backward") + ) + + +def _segment_tiles(index: torch.Tensor, num_rows: int): + """Sorted positions and fixed-size tiles per row segment (all integer ops).""" + + sorted_pos = torch.argsort(index, stable=True) + counts = torch.bincount(index, minlength=num_rows) + seg_start = torch.cumsum(counts, 0) - counts + tiles_per_seg = (counts + BACKWARD_TILE - 1) // BACKWARD_TILE + seg_first_tile = torch.zeros(num_rows + 1, dtype=torch.int64, device=index.device) + seg_first_tile[1:] = torch.cumsum(tiles_per_seg, 0) + tile_seg = torch.repeat_interleave(torch.arange(num_rows, device=index.device), tiles_per_seg) + tile_local = torch.arange(tile_seg.numel(), device=index.device) - seg_first_tile[tile_seg] + tile_begin = seg_start[tile_seg] + tile_local * BACKWARD_TILE + tile_end = torch.minimum(tile_begin + BACKWARD_TILE, seg_start[tile_seg] + counts[tile_seg]) + return sorted_pos, tile_begin.contiguous(), tile_end.contiguous(), seg_first_tile + + +class _H3AdaLNRowGatherCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, rows, timestep_indices, token_tags): + out = _C.h3_adaln_row_gather_forward( + rows, timestep_indices, token_tags, H3_ADALN_CHUNKS, H3_MODALITY_NUM + ) + ctx.save_for_backward(timestep_indices, token_tags) + ctx.rows_meta = (rows.shape[0], rows.dtype) + return out + + @staticmethod + def backward(ctx, grad_out): + timestep_indices, token_tags = ctx.saved_tensors + num_rows, rows_dtype = ctx.rows_meta + index = h3_adaln_indices(timestep_indices, token_tags) + tiles = _segment_tiles(index, num_rows) + grad_rows = _C.h3_adaln_row_gather_backward(grad_out.contiguous(), *tiles, rows_dtype) + record_backward( + "adaln_row_gather", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return grad_rows, None, None + + +class H3AdaLNRowGatherCudaOp: + """CUDA candidate: one launch gathers all six modulation tensors. + + The forward is a byte copy (bitwise equal to six ``index_select`` calls). + The backward replaces ``index_select``'s atomic BF16 scatter-add with a + deterministic FP32 segmented sum in stable sorted order, cast once. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + if not adaln_row_gather_available(): + raise RuntimeError( + "rl_engine._C lacks h3_adaln_row_gather_*; rebuild with " + "csrc/cuda/h3/adaln_row_gather.cu" + ) + + def __call__(self, rows, timestep_indices, token_tags): + return self.forward(rows, timestep_indices, token_tags) + + def forward(self, rows, timestep_indices, token_tags, *, check_range: bool = True): + validate_h3_adaln_row_gather(rows, timestep_indices, token_tags, check_range=check_range) + if not rows.is_cuda: + raise ValueError("H3AdaLNRowGatherCudaOp needs CUDA tensors") + if rows.stride(1) != 1: + rows = rows.contiguous() + packed = _H3AdaLNRowGatherCuda.apply( + rows, timestep_indices.contiguous(), token_tags.contiguous() + ) + return tuple(packed.unbind(0)) + + def gather_chunks(self, chunks, timestep_indices, token_tags, **kwargs): + """Drop-in for diffusers' six ``(3T, H)`` tensors: concatenate, then gather.""" + + if len(chunks) != H3_ADALN_CHUNKS: + raise ValueError(f"expected {H3_ADALN_CHUNKS} modulation tensors, got {len(chunks)}") + return self.forward(torch.cat(list(chunks), dim=1), timestep_indices, token_tags, **kwargs) + + def forward_fp32(self, rows, timestep_indices, token_tags): + return tuple(out.float() for out in self.forward(rows, timestep_indices, token_tags)) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/det_linear.py b/rl_engine/backends/cuda/model_specific/minimax_h3/det_linear.py new file mode 100644 index 000000000..1274364f2 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/det_linear.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Python surface of the H3 deterministic row-wise linear (contract h3-det-linear-v1). + +See ``csrc/cuda/h3/det_linear.cu`` for the reduction order. Every helper +raises instead of falling back to cuBLAS when the extension is missing. +""" + +from __future__ import annotations + +import torch + +from rl_engine.backends.extension import _C, _EXT_AVAILABLE + +ACT_NONE = 0 +ACT_SILU = 1 +CONTRACT = "h3-det-linear-v1" +DINPUT_CHUNK = 64 # N rows per d_input partial (kDInputChunk) +_SYMBOLS = ( + "h3_det_linear_forward", + "h3_det_linear_backward_input", + "h3_det_linear_backward_weight", + "h3_det_linear_backward_input_partials", + "h3_det_linear_fold_chunks", +) + + +def det_linear_available() -> bool: + """Return whether the extension exposes all three H3 deterministic linear APIs.""" + + return bool(_EXT_AVAILABLE and all(hasattr(_C, name) for name in _SYMBOLS)) + + +def _require() -> None: + """Raise when a deterministic linear entry point is unavailable.""" + + if not det_linear_available(): + raise RuntimeError( + "rl_engine._C lacks the h3_det_linear_* symbols; rebuild the CUDA extension " + "with csrc/cuda/h3/det_linear.cu (no cuBLAS fallback is allowed)" + ) + + +def det_linear_forward( + x: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor | None, + *, + activation: int = ACT_NONE, + save_pre_activation: bool = False, +) -> list[torch.Tensor]: + """Project CUDA ``x`` of shape ``(T, K)`` with ``(N, K)`` weights and optional bias. + + Inputs share an FP32 or BF16 dtype/device, ``T > 0``, and ``K`` supports + 16-byte vector loads; BF16 forward needs SM80 or newer. Return ``[out]`` + in the input dtype, plus an FP32 pre-activation when requested. Activation + is identity or SiLU; contiguous copies preserve the reduction contract. + """ + + _require() + return _C.h3_det_linear_forward( + x.contiguous(), + weight.contiguous(), + None if bias is None else bias.contiguous(), + int(activation), + bool(save_pre_activation), + ) + + +def det_linear_backward_input( + grad: torch.Tensor, weight: torch.Tensor, out_dtype: torch.dtype +) -> torch.Tensor: + """Compute a deterministic ``(T, K)`` input VJP on the inputs' CUDA device. + + Convert nonempty ``(T, N)`` gradients to FP32 and combine them with FP32 + or BF16 ``(N, K)`` weights; return FP32 or BF16 as specified by ``out_dtype``. + """ + + _require() + return _C.h3_det_linear_backward_input( + grad.float().contiguous(), weight.contiguous(), out_dtype + ) + + +def det_linear_backward_input_partials(grad: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + """``(ceil(N / 64), T, K)`` FP32 chunk sums of ``grad @ weight``, before the fold.""" + + _require() + return _C.h3_det_linear_backward_input_partials(grad.float().contiguous(), weight.contiguous()) + + +def det_linear_fold_chunks(partial: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + """Ascending left fold of the chunk partials, cast once (the second half of d_input).""" + + _require() + return _C.h3_det_linear_fold_chunks(partial.contiguous(), out_dtype) + + +def det_linear_backward_weight( + grad: torch.Tensor, x: torch.Tensor, w_dtype: torch.dtype, *, with_bias: bool = True +) -> list[torch.Tensor]: + """Fold nonempty CUDA ``(T, N)`` gradients and ``(T, K)`` inputs in row order. + + Convert gradients to FP32; ``x`` is FP32 or BF16 on the same device. + Return ``[dweight]`` of shape ``(N, K)``, optionally followed by ``dbias`` + of shape ``(N,)``, both in the requested FP32 or BF16 ``w_dtype``. + """ + + _require() + return _C.h3_det_linear_backward_weight( + grad.float().contiguous(), x.contiguous(), w_dtype, bool(with_bias) + ) + + +def silu_backward_fp32(grad: torch.Tensor, pre_activation: torch.Tensor) -> torch.Tensor: + """Elementwise SiLU VJP in FP32: g * s * (1 + z * (1 - s)), s = sigmoid(z).""" + + z = pre_activation.float() + s = torch.sigmoid(z) + return grad.float() * (s * (1.0 + z * (1.0 - s))) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/final_adaln_out.py b/rl_engine/backends/cuda/model_specific/minimax_h3/final_adaln_out.py new file mode 100644 index 000000000..2ff1fd16a --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/final_adaln_out.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 final AdaLN output (RFC #420 ``final_adaln_out``).""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C +from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import _segment_tiles +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + det_linear_available, + det_linear_backward_input, + det_linear_backward_weight, + det_linear_forward, + silu_backward_fp32, +) +from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import h3_rmsnorm_available +from rl_engine.reference.minimax_h3.final_adaln_out import validate_h3_final_adaln_out +from rl_engine.reference.minimax_h3.rmsnorm import H3_NORM_EPS + +KERNEL_ID = ( + "torch.silu[fp32]+cast[weight_dtype]+rl_engine._C.h3_det_linear_forward" + "+rl_engine._C.h3_rmsnorm_forward[modulated]" +) +BACKWARD_IMPL = "fp32_table_grad_into_h3_det_linear_backward" +BACKWARD_KERNEL_ID = ( + "rl_engine._C.h3_rmsnorm_backward+rl_engine._C.h3_det_linear_backward_weight" + "+rl_engine._C.h3_det_linear_backward_input" + "+rl_engine.backends.cuda.model_specific.minimax_h3.det_linear.silu_backward_fp32" +) + + +class _H3FinalAdaLNOutCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, x, norm_weight, temb, weight, bias, timestep_indices, eps): + hidden = x.shape[-1] + temb = temb.contiguous() + act = F.silu(temb).to(weight.dtype) # declared boundary: one cast after the FP32 SiLU + (table,) = det_linear_forward(act, weight, bias) + shift, scale = table.chunk(2, dim=-1) + x2 = x.contiguous().view(-1, hidden) + y, rstd = _C.h3_rmsnorm_forward( + x2, norm_weight.contiguous(), float(eps), shift, scale, timestep_indices + ) + ctx.save_for_backward(x2, norm_weight, rstd, temb, act, weight, table, timestep_indices) + ctx.x_shape = x.shape + return y.view(x.shape) + + @staticmethod + def backward(ctx, grad): + x2, norm_weight, rstd, temb, act, weight, table, timestep_indices = ctx.saved_tensors + shift, scale = table.chunk(2, dim=-1) + positions = timestep_indices.repeat(x2.shape[0] // timestep_indices.shape[0]) + tiles = _segment_tiles(positions, table.shape[0]) + dx, d_norm_weight, d_shift, d_scale = _C.h3_rmsnorm_backward( + grad.contiguous().view_as(x2), + x2, + norm_weight.contiguous(), + rstd, + shift, + scale, + timestep_indices, + *tiles, + ) + # The table gradient stays FP32 into the projection backward. + d_table = torch.cat([d_shift, d_scale], dim=1) + dw, db = det_linear_backward_weight(d_table, act, weight.dtype) + d_temb = silu_backward_fp32(det_linear_backward_input(d_table, weight, torch.float32), temb) + record_backward( + "final_adaln_out", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return dx.view(ctx.x_shape), d_norm_weight, d_temb, dw, db, None, None + + +class H3FinalAdaLNOutCudaOp: + """CUDA candidate: tensor-core shift/scale projection + fused norm/modulation. + + The projection is ``adaln_projection_3mod``'s deterministic GEMV + (``h3-det-linear-bf16-mma-v1``) on ``norm_out.linear``; the norm and + modulation are ``h3_rmsnorm``'s kernel indexed by ``timestep_indices``. + One autograd node, so the table gradient reaches the projection in FP32. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + if not (det_linear_available() and h3_rmsnorm_available()): + raise RuntimeError("rl_engine._C lacks the h3_det_linear_* / h3_rmsnorm_* ops") + + def __call__(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + return self.forward(x, norm_weight, temb, weight, bias, timestep_indices, eps) + + def forward(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + validate_h3_final_adaln_out(x, norm_weight, temb, weight, bias, timestep_indices, eps) + if weight.dtype != x.dtype: + raise TypeError("norm_out.linear must share the activations' dtype") + if not x.is_cuda: + raise ValueError("H3FinalAdaLNOutCudaOp needs CUDA tensors") + return _H3FinalAdaLNOutCuda.apply( + x, norm_weight, temb, weight, bias, timestep_indices.contiguous(), eps + ) + + def forward_fp32(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + return self.forward(x, norm_weight, temb, weight, bias, timestep_indices, eps).float() diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/gate_residual.py b/rl_engine/backends/cuda/model_specific/minimax_h3/gate_residual.py new file mode 100644 index 000000000..2c6c6be55 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/gate_residual.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 gated residual (RFC #420 ``adaln_gate_residual``).""" + +from __future__ import annotations + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C, _EXT_AVAILABLE +from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import _segment_tiles +from rl_engine.reference.minimax_h3.gate_residual import validate_h3_gate_residual + +KERNEL_ID = "rl_engine._C.h3_gate_residual_forward" +BACKWARD_IMPL = "exact_dy_sorted_segment_dgate" +BACKWARD_KERNEL_ID = "torch.argsort[stable]+rl_engine._C.h3_gate_residual_backward" + + +def h3_gate_residual_available() -> bool: + return bool( + _EXT_AVAILABLE + and hasattr(_C, "h3_gate_residual_forward") + and hasattr(_C, "h3_gate_residual_backward") + ) + + +class _H3GateResidualCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, residual, y, gate, index): + hidden = residual.shape[-1] + res2 = residual.contiguous().view(-1, hidden) + y2 = y.contiguous().view(-1, hidden) + out = _C.h3_gate_residual_forward(res2, y2, gate, index) + ctx.save_for_backward(y2, gate, index) + ctx.shape = residual.shape + return out.view(residual.shape) + + @staticmethod + def backward(ctx, grad): + y2, gate, index = ctx.saved_tensors + g2 = grad.contiguous().view_as(y2) + positions = index.repeat(y2.shape[0] // index.shape[0]) + tiles = _segment_tiles(positions, gate.shape[0]) + dy, dgate = _C.h3_gate_residual_backward(g2, y2, gate, index, *tiles) + record_backward( + "adaln_gate_residual", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return grad, dy.view(ctx.shape), dgate, None + + +class H3GateResidualCudaOp: + """CUDA candidate: ``residual + gate[index] * y`` in one elementwise pass. + + The gate row is gathered in the kernel from the AdaLN table view, and the + two roundings match the eager expression, so the output is bitwise equal + to diffusers. ``dy`` is bitwise equal to the eager VJP; ``dgate`` is a + deterministic FP32 segment sum instead of ``index_select``'s BF16 atomics. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + if not h3_gate_residual_available(): + raise RuntimeError( + "rl_engine._C lacks h3_gate_residual_*; rebuild with csrc/cuda/h3/gate_residual.cu" + ) + + def __call__(self, residual, y, gate, index): + return self.forward(residual, y, gate, index) + + def forward(self, residual, y, gate, index, *, check_range: bool = True): + validate_h3_gate_residual(residual, y, gate, index, check_range=check_range) + if not residual.is_cuda: + raise ValueError("H3GateResidualCudaOp needs CUDA tensors") + if gate.stride(1) != 1: + gate = gate.contiguous() + return _H3GateResidualCuda.apply(residual, y, gate, index.contiguous()) + + def forward_fp32(self, residual, y, gate, index): + return self.forward(residual, y, gate, index).float() diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/rmsnorm.py b/rl_engine/backends/cuda/model_specific/minimax_h3/rmsnorm.py new file mode 100644 index 000000000..a6ed3a28b --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/rmsnorm.py @@ -0,0 +1,103 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 RMSNorm with fused AdaLN modulation (RFC #420 ``h3_rmsnorm``).""" + +from __future__ import annotations + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C, _EXT_AVAILABLE +from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import _segment_tiles +from rl_engine.reference.minimax_h3.rmsnorm import ( + H3_NORM_EPS, + validate_h3_modulation, + validate_h3_rmsnorm, +) + +KERNEL_ID = "rl_engine._C.h3_rmsnorm_forward" +BACKWARD_IMPL = "row_local_dx_tiled_dweight_sorted_table_grads" +BACKWARD_KERNEL_ID = "torch.argsort[stable]+rl_engine._C.h3_rmsnorm_backward" + + +def h3_rmsnorm_available() -> bool: + return bool( + _EXT_AVAILABLE and hasattr(_C, "h3_rmsnorm_forward") and hasattr(_C, "h3_rmsnorm_backward") + ) + + +class _H3RMSNormCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight, shift, scale, index, eps): + hidden = x.shape[-1] + x2 = x.contiguous().view(-1, hidden) + y, rstd = _C.h3_rmsnorm_forward(x2, weight.contiguous(), float(eps), shift, scale, index) + ctx.save_for_backward(x2, weight, rstd, shift, scale, index) + ctx.x_shape = x.shape + ctx.modulated = index is not None + return y.view(x.shape) + + @staticmethod + def backward(ctx, grad): + x2, weight, rstd, shift, scale, index = ctx.saved_tensors + g2 = grad.contiguous().view_as(x2) + if not ctx.modulated: + dx, dweight = _C.h3_rmsnorm_backward(g2, x2, weight.contiguous(), rstd) + d_shift = d_scale = None + else: + positions = index.repeat(x2.shape[0] // index.shape[0]) + tiles = _segment_tiles(positions, shift.shape[0]) + dx, dweight, d_shift, d_scale = _C.h3_rmsnorm_backward( + g2, x2, weight.contiguous(), rstd, shift, scale, index, *tiles + ) + d_shift, d_scale = d_shift.to(shift.dtype), d_scale.to(scale.dtype) + record_backward( + "h3_rmsnorm", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return dx.view(ctx.x_shape), dweight, d_shift, d_scale, None, None + + +class H3RMSNormCudaOp: + """CUDA candidate. + + The statistics replay PyTorch's own RMSNorm reduction order, so the plain + norm is bitwise equal to ``nn.RMSNorm``; the modulated form gathers its + ``shift``/``scale`` rows in the kernel (no ``(S, H)`` intermediates) and + rounds each step where the eager expression does, so it is bitwise equal to + diffusers' ``n * (1.0 + scale[i]) + shift[i]`` as well. Rows are + independent of batch size and position. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + if not h3_rmsnorm_available(): + raise RuntimeError( + "rl_engine._C lacks h3_rmsnorm_*; rebuild with csrc/cuda/h3/rmsnorm_modulate.cu" + ) + + def __call__(self, x, weight, eps: float = H3_NORM_EPS): + return self.forward(x, weight, eps) + + def forward(self, x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + validate_h3_rmsnorm(x, weight, eps) + if not x.is_cuda: + raise ValueError("H3RMSNormCudaOp needs CUDA tensors") + return _H3RMSNormCuda.apply(x, weight, None, None, None, eps) + + def forward_modulated( + self, x, weight, shift, scale, index, eps: float = H3_NORM_EPS, *, check_range=True + ) -> torch.Tensor: + validate_h3_rmsnorm(x, weight, eps) + validate_h3_modulation(x, shift, scale, index, check_range=check_range) + if not x.is_cuda: + raise ValueError("H3RMSNormCudaOp needs CUDA tensors") + if shift.stride(1) != 1 or scale.stride(1) != 1 or shift.stride(0) != scale.stride(0): + shift, scale = shift.contiguous(), scale.contiguous() + return _H3RMSNormCuda.apply(x, weight, shift, scale, index.contiguous(), eps) + + def forward_fp32(self, x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + return self.forward(x, weight, eps).float() diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/sp_norm_adaln.py b/rl_engine/backends/cuda/model_specific/minimax_h3/sp_norm_adaln.py new file mode 100644 index 000000000..adf01094e --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/sp_norm_adaln.py @@ -0,0 +1,414 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Sequence-parallel H3 RMSNorm / AdaLN modulation / gated residual (RFC #420 ``sp_norm_adaln``). + +Ownership: rank ``r`` of ``sp`` holds packed positions ``[r * S // sp, (r + 1) * S // sp)`` +of every batch item. The full ``(S,)`` row index is replicated; it is small, and every +rank derives the same reduction plan from it. + +Forward and the row-local gradients (``dx``, ``d_sublayer``, ``d_residual``) are the WS1 +kernels on the local rows. Rows are independent, so they are byte-equal by construction. + +The cross-row reductions (``d_norm_weight``, ``d_shift``/``d_scale``, ``d_gate``) are WS1 +two-level folds: FP32 sums over fixed 256-element tiles, then an ascending fold of the +tiles. Every rank builds the global WS1 tile list. A tile is computed by the rank that +holds its first row; the tile's rows from other ranks are all-gathered beforehand (only +rows of tiles that straddle a shard boundary move). The tile partials are all-gathered, +put back in WS1 tile order, and every rank runs the WS1 fold. The partial and fold kernels +are the WS1 kernels, so every gradient equals WS1 on every rank. The only collectives are +rank-ordered all-gathers (copies). +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C, _EXT_AVAILABLE +from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + BACKWARD_TILE, + _segment_tiles, +) +from rl_engine.backends.cuda.model_specific.minimax_h3.ws2_comm import gather_rows +from rl_engine.reference.minimax_h3.gate_residual import validate_h3_gate_residual +from rl_engine.reference.minimax_h3.rmsnorm import ( + H3_NORM_EPS, + validate_h3_modulation, + validate_h3_rmsnorm, +) + +_SYMBOLS = ( + "h3_rmsnorm_forward", + "h3_rmsnorm_backward_dx", + "h3_rmsnorm_backward_partials", + "h3_rmsnorm_fold_partials", + "h3_gate_residual_forward", + "h3_gate_residual_backward_dy", + "h3_gate_grad_partials", + "h3_gate_grad_fold", +) +BACKWARD_IMPL = "row_local_ws1+gather_straddling_rows+ws1_tile_partials+gather+ws1_fold" + + +def sp_norm_adaln_available() -> bool: + return bool(_EXT_AVAILABLE and all(hasattr(_C, name) for name in _SYMBOLS)) + + +@dataclass(frozen=True) +class SPRowLayout: + """Which packed positions each sequence-parallel rank holds.""" + + seq_len: int + batch: int + sp: int + rank: int + + def bounds(self, rank: int) -> tuple[int, int]: + return rank * self.seq_len // self.sp, (rank + 1) * self.seq_len // self.sp + + @property + def lo(self) -> int: + return self.bounds(self.rank)[0] + + @property + def hi(self) -> int: + return self.bounds(self.rank)[1] + + @property + def local_len(self) -> int: + return self.hi - self.lo + + +def sp_row_layout(seq_len: int, batch: int, sp: int, rank: int) -> SPRowLayout: + if sp < 1 or not 0 <= rank < sp: + raise ValueError(f"need 0 <= rank < sp, got rank={rank}, sp={sp}") + if batch < 1 or seq_len < sp: + raise ValueError(f"every rank needs a row: seq_len={seq_len}, sp={sp}, batch={batch}") + return SPRowLayout(seq_len, batch, sp, rank) + + +def _ordinal_by_group(group: torch.Tensor, groups: int) -> tuple[torch.Tensor, torch.Tensor]: + """Per-group counts, and each element's position within its group (in input order).""" + + counts = torch.bincount(group, minlength=groups) + order = torch.argsort(group, stable=True) + ordinal = torch.empty_like(group) + first = torch.cumsum(counts, 0) - counts + ordinal[order] = torch.arange(group.numel(), device=group.device) - first[group[order]] + return counts, ordinal + + +class _Plan: + """Global WS1 tiles of one backward, and this rank's part in computing them. + + ``families`` are ``(rows, tile_begin, tile_end)`` over global flattened rows + ``b * S + s``; consecutive tiles cover ``rows`` in order. Every rank builds the same + plan from replicated metadata, so collective calls and shapes agree. + """ + + def __init__(self, layout: SPRowLayout, families, device) -> None: + self.layout, sp = layout, layout.sp + self._his = torch.tensor([layout.bounds(r)[1] for r in range(sp)], device=device) + + # Rows of a tile computed by another rank: the only activation rows that move. + tile_owner, shared = [], [] + for rows, begin, end in families: + owner = self.owner(rows[begin]) + tile_owner.append(owner) + shared.append(rows[torch.repeat_interleave(owner, end - begin) != self.owner(rows)]) + shared = torch.unique(torch.cat(shared)) + shared_owner = self.owner(shared) + counts, slot = _ordinal_by_group(shared_owner, sp) + self.max_send = int(counts.max()) if shared.numel() else 0 + + # Buffer row of every global row this rank reads: its own rows first, then + # every rank's send block (padded to max_send) in rank order. + own = self.global_rows(layout.rank) + buffer_index = torch.full((layout.batch * layout.seq_len,), -1, device=device) + buffer_index[own] = torch.arange(own.numel(), device=device) + foreign = shared_owner != layout.rank + buffer_index[shared[foreign]] = ( + own.numel() + shared_owner[foreign] * self.max_send + slot[foreign] + ) + mine = ~foreign + self.send_local = self.local_index(shared[mine][torch.argsort(slot[mine])]) + self.recv_global = torch.zeros(sp * self.max_send, dtype=torch.long, device=device) + self.recv_global[shared_owner * self.max_send + slot] = shared # padding: never read + + self.families = [ + self._my_tiles(rows, begin, end, owner, buffer_index) + for (rows, begin, end), owner in zip(families, tile_owner) + ] + + def owner(self, rows: torch.Tensor) -> torch.Tensor: + return torch.searchsorted(self._his, rows % self.layout.seq_len, right=True) + + def global_rows(self, rank: int) -> torch.Tensor: + lo, hi = self.layout.bounds(rank) + s = torch.arange(lo, hi, device=self._his.device) + b = torch.arange(self.layout.batch, device=self._his.device) + return (b[:, None] * self.layout.seq_len + s[None, :]).reshape(-1) + + def local_index(self, rows: torch.Tensor) -> torch.Tensor: + b, s = rows // self.layout.seq_len, rows % self.layout.seq_len + return b * self.layout.local_len + s - self.layout.lo + + def _my_tiles(self, rows, begin, end, owner, buffer_index): + mine = torch.nonzero(owner == self.layout.rank).flatten() + lengths = (end - begin)[mine] + my_end = torch.cumsum(lengths, 0) + my_begin = my_end - lengths + elems = torch.repeat_interleave(begin[mine] - my_begin, lengths) + elems = elems + torch.arange(elems.numel(), device=rows.device) + counts, ordinal = _ordinal_by_group(owner, self.layout.sp) + width = int(counts.max()) if owner.numel() else 0 + return { + "tiles": (buffer_index[rows[elems]].contiguous(), my_begin, my_end), + "width": width, # tiles per rank in the padded partial gather + "gather_index": owner * width + ordinal, # WS1 tile order within that gather + } + + def exchange(self, collective, *tensors: torch.Tensor) -> list[torch.Tensor]: + """For each ``(M_local, ...)`` tensor: local rows, then the gathered send blocks.""" + + out = [] + for t in tensors: + if self.max_send == 0: + out.append(t) + continue + send = t.new_zeros((self.max_send, *t.shape[1:])) + send[: self.send_local.numel()] = t[self.send_local] + out.append(torch.cat([t, gather_rows(collective, send)])) + return out + + def gather_partials(self, collective, family: int, partial: torch.Tensor) -> torch.Tensor: + """This rank's tile partials -> every tile's partial, in WS1 tile order.""" + + fam = self.families[family] + padded = partial.new_zeros((fam["width"], *partial.shape[1:])) + padded[: partial.shape[0]] = partial + return gather_rows(collective, padded).index_select(0, fam["gather_index"]).contiguous() + + +def _dweight_tiles(total: int, device): + begin = torch.arange(0, total, BACKWARD_TILE, device=device) + return torch.arange(total, device=device), begin, torch.clamp(begin + BACKWARD_TILE, max=total) + + +@dataclass(frozen=True) +class SPPlan: + """Everything an SP backward needs that depends only on the layout and the row index. + + Built once per ``(row index, table rows)`` and shared by every norm and gated + residual that uses that index (every block of a forward pass). + """ + + layout: SPRowLayout + plan: _Plan + local_index: torch.Tensor | None # (S_local,) table row of each local position + index_buf: torch.Tensor | None # table row of each exchange-buffer row + seg_first_tile: torch.Tensor | None + num_rows: int | None + + +def sp_plan( + layout: SPRowLayout, index_full=None, num_rows: int | None = None, *, device=None +) -> SPPlan: + """The reduction plan for one row index (``None``: the plain norm's ``dweight`` only).""" + + device = index_full.device if index_full is not None else torch.device(device or "cuda") + families = [_dweight_tiles(layout.batch * layout.seq_len, device)] + if index_full is None: + return SPPlan(layout, _Plan(layout, families, device), None, None, None, None) + *seg_tiles, seg_first_tile = _segment_tiles(index_full.repeat(layout.batch), num_rows) + plan = _Plan(layout, [*families, tuple(seg_tiles)], device) + local_index = index_full[layout.lo : layout.hi].clone(memory_format=torch.contiguous_format) + received = index_full[plan.recv_global % layout.seq_len] + index_buf = torch.cat([local_index.repeat(layout.batch), received]).contiguous() + return SPPlan(layout, plan, local_index, index_buf, seg_first_tile, num_rows) + + +def sp_rmsnorm_backward(grad, x, weight, rstd, shift, scale, sp: SPPlan, collective): + """One rank's backward on its ``(B * S_local, H)`` rows: ``(dx, dweight, d_shift, d_scale)``. + + ``dx`` is row-local; the other gradients are the full WS1 reductions, identical on + every rank. ``shift``/``scale`` are ``None`` for the plain norm. + """ + + weight, plan, modulated = weight.contiguous(), sp.plan, sp.local_index is not None + dx = _C.h3_rmsnorm_backward_dx(grad, x, weight, rstd, shift, scale, sp.local_index) + g_buf, x_buf, rstd_buf = plan.exchange(collective, grad, x, rstd) + seg = plan.families[1]["tiles"] if modulated else (None, None, None) + partials = _C.h3_rmsnorm_backward_partials( + g_buf, x_buf, weight, rstd_buf, shift, scale, sp.index_buf, + *plan.families[0]["tiles"], *seg, + ) # fmt: skip + gathered = [plan.gather_partials(collective, i, p) for i, p in enumerate(partials)] + record_backward( + "sp_norm_adaln", + kernel_id="rl_engine._C.h3_rmsnorm_backward_partials", + impl=BACKWARD_IMPL, + family="cuda", + ) + if not modulated: + (dweight,) = _C.h3_rmsnorm_fold_partials(gathered[0], weight) + return dx, dweight, None, None + dweight, d_shift, d_scale = _C.h3_rmsnorm_fold_partials( + gathered[0], weight, gathered[1], sp.seg_first_tile + ) + return dx, dweight, d_shift.to(shift.dtype), d_scale.to(scale.dtype) + + +def sp_gate_residual_backward(grad, y, gate, sp: SPPlan, collective): + """One rank's backward on its rows: ``(d_sublayer, d_gate)``; ``d_residual`` is ``grad``.""" + + plan = sp.plan + dy = _C.h3_gate_residual_backward_dy(grad, gate, sp.local_index) + g_buf, y_buf = plan.exchange(collective, grad, y) + partial = _C.h3_gate_grad_partials(g_buf, y_buf, *plan.families[1]["tiles"]) + dgate = _C.h3_gate_grad_fold( + plan.gather_partials(collective, 1, partial), sp.seg_first_tile, gate.dtype + ) + record_backward( + "sp_norm_adaln", + kernel_id="rl_engine._C.h3_gate_grad_partials", + impl=BACKWARD_IMPL, + family="cuda", + ) + return dy, dgate + + +class _SPRMSNorm(torch.autograd.Function): + @staticmethod + def forward(ctx, x, weight, shift, scale, sp, collective, eps): + x2 = x.contiguous().view(-1, x.shape[-1]) + y, rstd = _C.h3_rmsnorm_forward( + x2, weight.contiguous(), float(eps), shift, scale, sp.local_index + ) + ctx.save_for_backward(x2, weight, rstd, shift, scale) + ctx.sp, ctx.collective, ctx.x_shape = sp, collective, x.shape + return y.view(x.shape) + + @staticmethod + def backward(ctx, grad): + x2, weight, rstd, shift, scale = ctx.saved_tensors + g2 = grad.contiguous().view_as(x2) + dx, dweight, d_shift, d_scale = sp_rmsnorm_backward( + g2, x2, weight, rstd, shift, scale, ctx.sp, ctx.collective + ) + return dx.view(ctx.x_shape), dweight, d_shift, d_scale, None, None, None + + +class _SPGateResidual(torch.autograd.Function): + @staticmethod + def forward(ctx, residual, y, gate, sp, collective): + hidden = residual.shape[-1] + res2 = residual.contiguous().view(-1, hidden) + y2 = y.contiguous().view(-1, hidden) + out = _C.h3_gate_residual_forward(res2, y2, gate, sp.local_index) + ctx.save_for_backward(y2, gate) + ctx.sp, ctx.collective, ctx.shape = sp, collective, residual.shape + return out.view(residual.shape) + + @staticmethod + def backward(ctx, grad): + y2, gate = ctx.saved_tensors + g2 = grad.contiguous().view_as(y2) + dy, dgate = sp_gate_residual_backward(g2, y2, gate, ctx.sp, ctx.collective) + return grad, dy.view(ctx.shape), dgate, None, None + + +class H3SPNormAdaLNCudaOp: + """Sequence-parallel ``h3_rmsnorm`` (+ modulation) and ``adaln_gate_residual``. + + ``collective`` is a rank-ordered all-gather with ``rank``, ``world_size`` and + ``backend_id`` (e.g. ``DeterministicCollective``). Activations are this rank's + ``(B, S_local, H)`` rows; the norm weight, the table views and the ``(S,)`` row + index are the full, replicated tensors. Gradients of the replicated tensors come + back identical on every rank, equal to WS1. + """ + + op_class = "reduction" + backward_impl = BACKWARD_IMPL + + def __init__(self, collective, seq_len: int, batch: int) -> None: + if not sp_norm_adaln_available(): + raise RuntimeError("rl_engine._C lacks the H3 SP symbols; rebuild the CUDA extension") + self.collective = collective + self.layout = sp_row_layout(seq_len, batch, collective.world_size, collective.rank) + self._plans: list[tuple[object, int, int | None, torch.device, SPPlan]] = [] + + def plan(self, index: torch.Tensor | None, num_rows: int | None, *, device=None) -> SPPlan: + """The cached plan for this index object (rebuilt if it was modified in place).""" + + device = index.device if index is not None else torch.device(device or "cuda") + if device.type == "cuda" and device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + version = -1 if index is None else index._version + for obj, ver, rows, cached_device, plan in self._plans: + if obj is index and ver == version and rows == num_rows and cached_device == device: + return plan + plan = sp_plan(self.layout, index, num_rows, device=device) + # Holding the index keeps its storage alive, so a new tensor can never alias it. + self._plans = [(index, version, num_rows, device, plan), *self._plans[:3]] + return plan + + def _check(self, x: torch.Tensor, index: torch.Tensor | None, num_rows: int | None) -> None: + expected = (self.layout.batch, self.layout.local_len) + if not isinstance(x, torch.Tensor) or x.dim() != 3 or tuple(x.shape[:2]) != expected: + shape = tuple(x.shape) if isinstance(x, torch.Tensor) else type(x).__name__ + raise ValueError(f"activations must be this rank's {expected} + (H,) rows, got {shape}") + if not x.is_cuda: + raise ValueError("H3SPNormAdaLNCudaOp needs CUDA tensors") + if index is None: + return + if index.dim() != 1 or index.shape[0] != self.layout.seq_len: + raise ValueError( + f"index must be the full ({self.layout.seq_len},) row index, " + f"got {tuple(index.shape)}" + ) + if index.numel() and (int(index.min()) < 0 or int(index.max()) >= num_rows): + raise IndexError(f"index must be in [0, {num_rows})") + + def norm(self, x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + self._check(x, None, None) + validate_h3_rmsnorm(x, weight, eps) + return _SPRMSNorm.apply( + x, weight, None, None, self.plan(None, None, device=x.device), self.collective, eps + ) + + def norm_modulated(self, x, weight, shift, scale, index, eps: float = H3_NORM_EPS): + self._check(x, index, shift.shape[0]) + validate_h3_rmsnorm(x, weight, eps) + validate_h3_modulation(x, shift, scale, index[self.layout.lo : self.layout.hi]) + if shift.stride(1) != 1 or scale.stride(1) != 1 or shift.stride(0) != scale.stride(0): + shift, scale = shift.contiguous(), scale.contiguous() + sp = self.plan(index, shift.shape[0]) + return _SPRMSNorm.apply(x, weight, shift, scale, sp, self.collective, eps) + + def gate_residual(self, residual, y, gate, index) -> torch.Tensor: + self._check(residual, index, gate.shape[0]) + validate_h3_gate_residual(residual, y, gate, index[self.layout.lo : self.layout.hi]) + if gate.stride(1) != 1: + gate = gate.contiguous() + sp = self.plan(index, gate.shape[0]) + return _SPGateResidual.apply(residual, y, gate, sp, self.collective) + + def readback(self) -> dict: + return { + "op": "sp_norm_adaln", + "backward_impl": BACKWARD_IMPL, + "collective_backend": getattr( + self.collective, "backend_id", type(self.collective).__name__ + ), + "collective_ops": ["all_gather"], + "reduction_order": "ws1_256_element_tiles_then_ascending_fold", + "sp": self.layout.sp, + "rank": self.layout.rank, + "positions": [self.layout.lo, self.layout.hi], + "batch": self.layout.batch, + "fallback": None, + } diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_mlp.py b/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_mlp.py new file mode 100644 index 000000000..aff81eaa8 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_mlp.py @@ -0,0 +1,116 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 FP32 timestep MLP (RFC #420 ``timestep_mlp_fp32``).""" + +from __future__ import annotations + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + ACT_NONE, + ACT_SILU, + CONTRACT, + det_linear_available, + det_linear_backward_input, + det_linear_backward_weight, + det_linear_forward, + silu_backward_fp32, +) +from rl_engine.reference.minimax_h3.timestep_mlp import validate_h3_timestep_mlp + +KERNEL_ID = "rl_engine._C.h3_det_linear_forward[silu]+rl_engine._C.h3_det_linear_forward" +BACKWARD_IMPL = "h3_det_linear_v1_chunked_dinput_ascending_row_dweight" +BACKWARD_KERNEL_ID = ( + "rl_engine._C.h3_det_linear_backward_weight" + "+rl_engine._C.h3_det_linear_backward_input" + "+rl_engine.backends.cuda.model_specific.minimax_h3.det_linear.silu_backward_fp32" +) + + +class _H3TimestepMLPCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, x, w1, b1, w2, b2): + """Evaluate the FP32 CUDA linear-SiLU-linear graph and save VJP inputs.""" + + x = x.contiguous() + hidden, pre = det_linear_forward(x, w1, b1, activation=ACT_SILU, save_pre_activation=True) + (out,) = det_linear_forward(hidden, w2, b2, activation=ACT_NONE) + ctx.save_for_backward(x, w1, w2, pre, hidden) + return out + + @staticmethod + def backward(ctx, grad_out): + """Return requested FP32 input and parameter VJPs with fixed reduction order.""" + + x, w1, w2, pre, hidden = ctx.saved_tensors + need_x, need_w1, need_b1, need_w2, need_b2 = ctx.needs_input_grad + grad_out = grad_out.float().contiguous() + dw2 = db2 = dw1 = db1 = dx = None + if need_w2 or need_b2: + dw2, db2 = det_linear_backward_weight(grad_out, hidden, torch.float32) + d_pre = None + if need_x or need_w1 or need_b1: + d_hidden = det_linear_backward_input(grad_out, w2, torch.float32) + d_pre = silu_backward_fp32(d_hidden, pre).contiguous() + if need_w1 or need_b1: + dw1, db1 = det_linear_backward_weight(d_pre, x, torch.float32) + if need_x: + dx = det_linear_backward_input(d_pre, w1, torch.float32) + record_backward( + "timestep_mlp_fp32", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return ( + dx, + dw1 if need_w1 else None, + db1 if need_b1 else None, + dw2 if need_w2 else None, + db2 if need_b2 else None, + ) + + +class H3TimestepMLPCudaOp: + """CUDA candidate: two warp-per-column GEMVs, SiLU fused into the first. + + The weights dominate the traffic (5.5 MB + 57.8 MB FP32), the rows are the + handful of distinct timesteps, and every output element follows contract + ``h3-det-linear-v1``, so a timestep's ``temb`` bytes do not depend on which + other timesteps share the call. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + contract = CONTRACT + + def __init__(self) -> None: + """Require all deterministic linear symbols in the CUDA extension.""" + + if not det_linear_available(): + raise RuntimeError( + "rl_engine._C lacks h3_det_linear_*; rebuild with csrc/cuda/h3/det_linear.cu" + ) + + def __call__(self, x, w1, b1, w2, b2): + """Return FP32 ``(T, D)`` timestep embeddings for same-device CUDA inputs.""" + + return self.forward(x, w1, b1, w2, b2) + + def forward(self, x, w1, b1, w2, b2) -> torch.Tensor: + """Apply two deterministic FP32 CUDA projections with an intervening SiLU. + + Require same-device FP32 tensors: ``x`` is nonempty ``(T, K)``, + weights are ``(H, K)`` and ``(D, H)``, and biases are ``(H,)`` and + ``(D,)``. Return ``(T, D)`` with autograd support for all five inputs. + """ + + validate_h3_timestep_mlp(x, w1, b1, w2, b2) + if not x.is_cuda: + raise ValueError("H3TimestepMLPCudaOp needs CUDA tensors") + return _H3TimestepMLPCuda.apply(x, w1, b1, w2, b2) + + def forward_fp32(self, x, w1, b1, w2, b2) -> torch.Tensor: + """Use the standard CUDA MLP, whose inputs and output are already FP32.""" + + return self.forward(x, w1, b1, w2, b2) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_sinusoid.py b/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_sinusoid.py new file mode 100644 index 000000000..4e9dfe8a4 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/timestep_sinusoid.py @@ -0,0 +1,116 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CUDA H3 sinusoidal timestep features (RFC #420 ``timestep_sinusoid_h3``).""" + +from __future__ import annotations + +import torch + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.extension import _C, _EXT_AVAILABLE +from rl_engine.reference.minimax_h3 import H3_FREQ_DIM, H3_MAX_PERIOD +from rl_engine.reference.minimax_h3.fixed_order import tree_sum_lastdim_fp32 +from rl_engine.reference.minimax_h3.timestep_sinusoid import ( + NativeH3TimestepSinusoidOp, + validate_h3_timesteps, +) + +KERNEL_ID = "rl_engine._C.h3_timestep_sinusoid_forward" +BACKWARD_IMPL = "row_local_fp32_analytic_tree_sum" + + +def h3_sinusoid_cuda_available() -> bool: + """Return whether the extension provides the CUDA H3 sinusoid entry point.""" + + return bool(_EXT_AVAILABLE and hasattr(_C, "h3_timestep_sinusoid_forward")) + + +class _H3TimestepSinusoidCuda(torch.autograd.Function): + @staticmethod + def forward(ctx, timestep: torch.Tensor, num_channels: int, check_range: bool): + """Compute FP32 CUDA features and save them for the analytic timestep VJP.""" + + t32 = timestep.detach().float().contiguous() + out = _C.h3_timestep_sinusoid_forward( + t32, int(num_channels), float(H3_MAX_PERIOD), check_range + ) + ctx.save_for_backward(out) + ctx.num_channels = int(num_channels) + ctx.timestep_dtype = timestep.dtype + return out + + @staticmethod + def backward(ctx, grad_out: torch.Tensor): + """Reduce channel derivatives with a fixed FP32 tree into the timestep dtype.""" + + (out,) = ctx.saved_tensors + half = ctx.num_channels // 2 + freq = NativeH3TimestepSinusoidOp.frequencies_fp32(ctx.num_channels, out.device) + cos_part, sin_part = out[:, :half], out[:, half:] + g = grad_out.float() + # d cos(t f)/dt = -f sin(t f); d sin(t f)/dt = f cos(t f). Row-local, + # summed over channels with a fixed pairwise tree. + contrib = torch.cat( + [-(g[:, :half] * sin_part) * freq, (g[:, half:] * cos_part) * freq], dim=-1 + ) + grad_t = tree_sum_lastdim_fp32(contrib).to(ctx.timestep_dtype) + record_backward( + "timestep_sinusoid_h3", + kernel_id="rl_engine.reference.minimax_h3.fixed_order.tree_sum_lastdim_fp32", + impl=BACKWARD_IMPL, + family="cuda", + ) + return grad_t, None, None + + +class H3TimestepSinusoidCudaOp: + """CUDA candidate: one thread per (timestep, frequency), FP32 output. + + Construction fails when the extension lacks the symbol, so the registry + falls back to the PyTorch reference instead of silently mis-dispatching. + """ + + op_class = "elementwise" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + + def __init__(self) -> None: + """Require the native H3 CUDA sinusoid symbol before creating this operator.""" + + if not h3_sinusoid_cuda_available(): + raise RuntimeError( + "rl_engine._C.h3_timestep_sinusoid_forward is unavailable; rebuild the " + "CUDA extension with csrc/cuda/h3/timestep_sinusoid.cu" + ) + + def __call__(self, timestep: torch.Tensor, *, num_channels: int = H3_FREQ_DIM): + """Return FP32 CUDA features for timesteps in ``[0, 1]`` with gradients.""" + + return self.forward(timestep, num_channels=num_channels) + + def forward( + self, + timestep: torch.Tensor, + *, + num_channels: int = H3_FREQ_DIM, + check_range: bool = True, + ) -> torch.Tensor: + """Return FP32 ``(T, num_channels)`` cosine-then-sine features on CUDA. + + Require nonempty 1-D FP32, FP16 or BF16 timesteps and a positive even + channel count. Native validation rejects nonfinite/out-of-range values + when ``check_range`` is true; backward returns the original input dtype. + """ + + # Native code owns value validation, avoiding a second host sync here. + validate_h3_timesteps(timestep, num_channels, check_range=False) + if not timestep.is_cuda: + raise ValueError("H3TimestepSinusoidCudaOp needs a CUDA timestep tensor") + return _H3TimestepSinusoidCuda.apply(timestep, num_channels, check_range) + + def forward_fp32(self, timestep: torch.Tensor, *, num_channels: int = H3_FREQ_DIM): + """Use the standard CUDA sinusoid path, which already returns FP32 features.""" + + # The op is FP32 end to end; the FP32 path is the op itself. + return self.forward(timestep, num_channels=num_channels) diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/tp_adaln_projection.py b/rl_engine/backends/cuda/model_specific/minimax_h3/tp_adaln_projection.py new file mode 100644 index 000000000..177ce97f1 --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/tp_adaln_projection.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Tensor-parallel H3 AdaLN projection (RFC #420 ``tp_adaln_3mod``). + +Ownership: rank ``r`` of ``tp`` owns the contiguous output columns +``[r * N / tp, (r + 1) * N / tp)`` of the ``N = 3 * 6 * H`` projection, i.e. the +same rows of ``adaln_proj.linear.weight``/``bias``. ``N / tp`` must be a +multiple of the 64-row ``d_input`` chunk of contract ``h3-det-linear-v1``. + +Every floating-point operation is the WS1 kernel itself, so every rank ends +with the WS1 bytes: + +* forward: each column depends only on ``x`` and its own weight row, so the + local GEMV produces the WS1 columns; a rank-ordered all-gather (a copy) + rebuilds the ``(T, N)`` table, and the six tensors and ``3T`` modality rows + are views of it exactly as in WS1; +* ``dW``/``db``: rows of the shard, computed locally (no cross-rank sum); +* ``d_temb``: each rank computes the WS1 chunk partials of its columns, the + all-gather puts them in global chunk order, and every rank runs the WS1 + ascending fold. No collective reduction arithmetic is used. + +The backward expects the table gradient to be identical on every TP rank +(the modulation consumes the full replicated table); it reads its own columns. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + +from rl_engine.ops.autograd.backward_runtime import record_backward +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + CONTRACT, + DINPUT_CHUNK, + det_linear_available, + det_linear_backward_input_partials, + det_linear_backward_weight, + det_linear_fold_chunks, + det_linear_forward, + silu_backward_fp32, +) +from rl_engine.backends.cuda.model_specific.minimax_h3.ws2_comm import gather_rows +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM +from rl_engine.reference.minimax_h3.adaln_projection import split_adaln_table + +KERNEL_ID = "rl_engine._C.h3_det_linear_forward[column_shard]+all_gather[rank_order]" +BACKWARD_IMPL = "local_dweight+all_gather_chunk_partials+ascending_fold" +BACKWARD_KERNEL_ID = ( + "rl_engine._C.h3_det_linear_backward_weight" + "+rl_engine._C.h3_det_linear_backward_input_partials" + "+all_gather+rl_engine._C.h3_det_linear_fold_chunks" +) + + +@dataclass(frozen=True) +class AdaLNColumnShard: + """The output columns a TP rank owns, and the WS1 objects they map to.""" + + tp: int + rank: int + n_total: int + begin: int + end: int + + @property + def hidden(self) -> int: + return self.n_total // (H3_MODALITY_NUM * H3_ADALN_CHUNKS) + + def slots(self) -> list[tuple[int, int, int, int]]: + """``(modality, chunk, h_begin, h_end)`` pieces covered by this shard.""" + + out, col = [], self.begin + while col < self.end: + slot, h0 = divmod(col, self.hidden) + h1 = min(self.hidden, h0 + self.end - col) + out.append((*divmod(slot, H3_ADALN_CHUNKS), h0, h1)) + col += h1 - h0 + return out + + +def adaln_column_shard(n_total: int, tp: int, rank: int) -> AdaLNColumnShard: + if tp < 1 or not 0 <= rank < tp: + raise ValueError(f"need 0 <= rank < tp, got rank={rank}, tp={tp}") + if n_total <= 0: + raise ValueError(f"N must be positive, got {n_total}") + if n_total % (H3_MODALITY_NUM * H3_ADALN_CHUNKS): + raise ValueError(f"N={n_total} is not 3 modalities x 6 chunks x H") + if n_total % (tp * DINPUT_CHUNK): + # A shard boundary inside a d_input chunk would split one WS1 partial. + raise ValueError( + f"tp={tp} does not split N={n_total} on {DINPUT_CHUNK}-column chunk boundaries" + ) + width = n_total // tp + return AdaLNColumnShard(tp, rank, n_total, rank * width, (rank + 1) * width) + + +def shard_adaln_projection(weight: torch.Tensor, bias: torch.Tensor, tp: int, rank: int): + """This rank's rows of the full projection weight and bias (views).""" + + shard = adaln_column_shard(weight.shape[0], tp, rank) + return weight[shard.begin : shard.end], bias[shard.begin : shard.end] + + +def _gather_columns(collective, local: torch.Tensor) -> torch.Tensor: + """``(T, N / tp)`` column shards -> the ``(T, N)`` table, in rank order (a copy).""" + + num_t, width = local.shape + rows = gather_rows(collective, local) # (tp * T, N / tp), rank-major + return rows.view(-1, num_t, width).transpose(0, 1).reshape(num_t, -1) + + +def _validate_shard_inputs(temb, weight, bias, shard: AdaLNColumnShard) -> None: + for name, tensor in (("temb", temb), ("weight", weight), ("bias", bias)): + if not isinstance(tensor, torch.Tensor): + raise TypeError(f"{name} must be a torch.Tensor") + if not tensor.is_cuda or tensor.device != temb.device: + raise ValueError(f"{name} must be on this rank's CUDA device, got {tensor.device}") + if temb.dtype != torch.float32: + raise TypeError(f"temb must be float32 (the SiLU runs before the cast), got {temb.dtype}") + if weight.dtype not in (torch.bfloat16, torch.float32) or bias.dtype != weight.dtype: + raise TypeError( + f"weight/bias must share bfloat16 or float32, got {weight.dtype}/{bias.dtype}" + ) + width = shard.end - shard.begin + if weight.dim() != 2 or weight.shape[0] != width or bias.shape != (width,): + raise ValueError( + f"rank {shard.rank} of {shard.tp} owns {width} rows, got weight " + f"{tuple(weight.shape)} and bias {tuple(bias.shape)}" + ) + if temb.dim() != 2 or temb.shape[0] == 0 or temb.shape[1] != weight.shape[1]: + raise ValueError( + f"temb must be a non-empty (T, {weight.shape[1]}) matrix, got {tuple(temb.shape)}" + ) + + +class _TPAdaLNProjection(torch.autograd.Function): + @staticmethod + def forward(ctx, temb, weight_shard, bias_shard, collective, shard): + temb = temb.contiguous() + act = F.silu(temb).to(weight_shard.dtype) + (local,) = det_linear_forward(act, weight_shard, bias_shard) + ctx.save_for_backward(temb, act, weight_shard) + ctx.collective, ctx.shard = collective, shard + return _gather_columns(collective, local) + + @staticmethod + def backward(ctx, grad_table): + temb, act, weight = ctx.saved_tensors + shard = ctx.shard + grad_local = grad_table[:, shard.begin : shard.end].float().contiguous() + dw, db = det_linear_backward_weight(grad_local, act, weight.dtype) + partial = det_linear_backward_input_partials(grad_local, weight) + d_act = det_linear_fold_chunks(gather_rows(ctx.collective, partial), torch.float32) + record_backward( + "tp_adaln_3mod", kernel_id=BACKWARD_KERNEL_ID, impl=BACKWARD_IMPL, family="cuda" + ) + return silu_backward_fp32(d_act, temb), dw, db, None, None + + +class H3TPAdaLNProjectionCudaOp: + """Column-parallel ``adaln_projection_3mod``, byte-equal to WS1 on every rank. + + ``collective`` is a rank-ordered all-gather with ``rank``, ``world_size`` + and ``backend_id`` (e.g. ``DeterministicCollective``). Inputs are the replicated + FP32 ``temb`` and this rank's weight/bias rows. + """ + + op_class = "reduction" + kernel_id = KERNEL_ID + backward_impl = BACKWARD_IMPL + contract = CONTRACT + + def __init__(self, collective, n_total: int) -> None: + if not det_linear_available(): + raise RuntimeError( + "rl_engine._C lacks h3_det_linear_*; rebuild with csrc/cuda/h3/det_linear.cu" + ) + self.collective = collective + self.shard = adaln_column_shard(n_total, collective.world_size, collective.rank) + + def __call__(self, temb, weight_shard, bias_shard): + return self.forward(temb, weight_shard, bias_shard) + + def forward(self, temb, weight_shard, bias_shard) -> tuple[torch.Tensor, ...]: + table = self.forward_table(temb, weight_shard, bias_shard) + return split_adaln_table(table, self.shard.hidden) + + def forward_table(self, temb, weight_shard, bias_shard) -> torch.Tensor: + _validate_shard_inputs(temb, weight_shard, bias_shard, self.shard) + return _TPAdaLNProjection.apply(temb, weight_shard, bias_shard, self.collective, self.shard) + + def readback(self) -> dict: + """Runtime identity of this rank's projection, for evidence and strict traces.""" + + return { + "op": "tp_adaln_3mod", + "kernel_id": KERNEL_ID, + "backward_kernel_id": BACKWARD_KERNEL_ID, + "contract": CONTRACT, + "collective_backend": getattr( + self.collective, "backend_id", type(self.collective).__name__ + ), + "collective_ops": ["all_gather"], + "reduction_order": "ws1_ascending_chunk_fold_after_rank_ordered_gather", + "tp": self.shard.tp, + "rank": self.shard.rank, + "columns": [self.shard.begin, self.shard.end], + "slots": [list(s) for s in self.shard.slots()], + "fallback": None, + } diff --git a/rl_engine/backends/cuda/model_specific/minimax_h3/ws2_comm.py b/rl_engine/backends/cuda/model_specific/minimax_h3/ws2_comm.py new file mode 100644 index 000000000..f03597b7d --- /dev/null +++ b/rl_engine/backends/cuda/model_specific/minimax_h3/ws2_comm.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Rank-ordered row gather shared by the H3 WS2 ops (a copy: no arithmetic).""" + +from __future__ import annotations + +import torch + +# DeterministicCollective sends all-gathers whose output is at most this many bytes +# through a single-block byte-copy kernel (csrc/cuda/distributed/ +# deterministic_collective.cu, kSingleBlockFastPathMaxBytes) that costs about 4 us per +# KiB on B200; larger outputs take the multi-block path (~45 us). Padding the shard +# past the threshold is cheaper for every non-trivial size. +_SINGLE_BLOCK_GATHER_MAX_BYTES = 256 * 1024 + + +def gather_rows(collective, shard: torch.Tensor) -> torch.Tensor: + """``(world * rows, ...)``: every rank's ``(rows, ...)`` shard, in rank order.""" + + shard = shard.contiguous() + rows, world = shard.shape[0], collective.world_size + row_bytes = shard[0].numel() * shard.element_size() if rows else 0 + if world == 1: + return shard + if row_bytes == 0: + return collective.all_gather(shard) + padded = max(rows, _SINGLE_BLOCK_GATHER_MAX_BYTES // (world * row_bytes) + 1) + if padded == rows: + return collective.all_gather(shard) + staged = torch.cat([shard, shard.new_zeros((padded - rows, *shard.shape[1:]))]) + out = collective.all_gather(staged).view(world, padded, *shard.shape[1:]) + return out[:, :rows].reshape(world * rows, *shard.shape[1:]) diff --git a/rl_engine/reference/minimax_h3/__init__.py b/rl_engine/reference/minimax_h3/__init__.py new file mode 100644 index 000000000..a01f9218b --- /dev/null +++ b/rl_engine/reference/minimax_h3/__init__.py @@ -0,0 +1,30 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""MiniMax-H3 (RFC #420) PyTorch goldens. + +Shape constants are pinned to ``MiniMaxAI/MiniMax-H3@42ed227`` (see +``rl_engine/validation/models/h3_manifest.json``). The ops accept other sizes so that +tiny synthetic shapes can be tested, but the H3 layout rules (three modality +rows per timestep, six modulation chunks) are fixed. +""" + +from __future__ import annotations + +H3_FREQ_DIM = 256 +H3_MAX_PERIOD = 10000 +H3_TIME_EMBED_HIDDEN_DIM = 5376 +H3_TIME_EMBED_DIM = 2688 +H3_HIDDEN_SIZE = 5376 +# Modality tags of the packed sequence: 0 video, 1 text, 2 audio. +H3_MODALITY_NUM = 3 +# shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp +H3_ADALN_CHUNKS = 6 +H3_ADALN_CHUNK_NAMES = ( + "shift_msa", + "scale_msa", + "gate_msa", + "shift_mlp", + "scale_mlp", + "gate_mlp", +) diff --git a/rl_engine/reference/minimax_h3/adaln_projection.py b/rl_engine/reference/minimax_h3/adaln_projection.py new file mode 100644 index 000000000..03c1ad683 --- /dev/null +++ b/rl_engine/reference/minimax_h3/adaln_projection.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 three-modality AdaLN projection (RFC #420 row ``adaln_projection_3mod``). + +One ``MiniMaxH3AdaLayerNormModulation`` per transformer block:: + + act = silu(temb).to(weight.dtype) SiLU in FP32, one cast to BF16 + table = act @ W.T + b (T, 6 * H * 3), BF16 + rows = table.view(3 * T, 6 * H) row t * 3 + m, m in {video, text, audio} + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = rows.chunk(6, -1) + +Output channel ``o = m * 6H + c * H + h`` of the projection is chunk ``c``, hidden +index ``h`` of modality ``m``. Each of the six outputs has shape ``(3T, H)``, +and row ``t * 3 + m`` is what ``timestep_indices * 3 + token_tags`` addresses. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM + +_WEIGHT_DTYPES = (torch.bfloat16, torch.float32) + + +def adaln_hidden_size(weight: torch.Tensor) -> int: + """Infer ``H`` from a 2-D projection weight whose row count is divisible by 18.""" + + rows_per_hidden = H3_ADALN_CHUNKS * H3_MODALITY_NUM + if weight.dim() != 2 or weight.shape[0] % rows_per_hidden != 0: + raise ValueError( + f"AdaLN weight must be (6 * H * 3, D); got {tuple(weight.shape)}, whose first " + f"dim is not a multiple of {rows_per_hidden}" + ) + return weight.shape[0] // rows_per_hidden + + +def validate_h3_adaln_projection(temb: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor): + """Check the AdaLN input shapes, shared device and precision boundary; return ``H``. + + Require nonempty FP32 ``temb`` of shape ``(T, D)`` and FP32 or BF16 weight + of shape ``(18H, D)`` with matching-dtype bias of shape ``(18H,)``. + """ + + for name, tensor in (("temb", temb), ("weight", weight), ("bias", bias)): + if not isinstance(tensor, torch.Tensor): + raise TypeError(f"{name} must be a torch.Tensor") + if tensor.device != temb.device: + raise ValueError(f"{name} is on {tensor.device}, temb is on {temb.device}") + if temb.dtype != torch.float32: + # RFC probe H7: the activation must run at the timestep-embedding + # precision; a BF16 temb means the cast happened before the SiLU. + raise TypeError(f"temb must be float32 (the SiLU runs before the cast), got {temb.dtype}") + if weight.dtype not in _WEIGHT_DTYPES or bias.dtype != weight.dtype: + raise TypeError( + f"weight/bias must share a dtype in {_WEIGHT_DTYPES}, got {weight.dtype}/{bias.dtype}" + ) + hidden = adaln_hidden_size(weight) + if temb.dim() != 2 or temb.shape[0] == 0 or temb.shape[1] != weight.shape[1]: + raise ValueError( + f"temb must be a non-empty (T, {weight.shape[1]}) matrix, got {tuple(temb.shape)}" + ) + if bias.shape != (weight.shape[0],): + raise ValueError(f"bias must be ({weight.shape[0]},), got {tuple(bias.shape)}") + return hidden + + +def split_adaln_table(table: torch.Tensor, hidden: int) -> tuple[torch.Tensor, ...]: + """(T, 6H*3) projection output -> six (3T, H) views, diffusers order.""" + + return table.view(-1, H3_ADALN_CHUNKS * hidden).chunk(H3_ADALN_CHUNKS, dim=-1) + + +class NativeH3AdaLNProjectionOp: + """PyTorch reference for the H3 AdaLN projection. + + ``forward`` is the provider path (``F.silu`` in FP32, cast, ``F.linear``). + ``forward_fp32`` is the golden: the SiLU in FP64, rounded once to the + weight dtype at the declared boundary (the cast is model semantics, not + an implementation detail), then the projection in FP64. The six outputs are + returned in FP32 without the final BF16 rounding. Its gradient is FP64 + end to end: the cast is applied straight-through (identity VJP). + """ + + op_class = "reduction" + + def __call__(self, temb, weight, bias): + """Return six weight-dtype ``(3T, H)`` views using the PyTorch provider path.""" + + return self.forward(temb, weight, bias) + + def forward(self, temb, weight, bias) -> tuple[torch.Tensor, ...]: + """Apply FP32 SiLU, cast to the weight dtype and project on the input device. + + Validate ``(T, D)`` embeddings and ``(18H, D)``/``(18H,)`` parameters; + return six weight-dtype ``(3T, H)`` views with PyTorch autograd support. + """ + + hidden = validate_h3_adaln_projection(temb, weight, bias) + table = F.linear(F.silu(temb).to(weight.dtype), weight, bias) + return split_adaln_table(table, hidden) + + def forward_fp32(self, temb, weight, bias) -> tuple[torch.Tensor, ...]: + """Return six FP32 golden outputs from FP64 SiLU and projection arithmetic. + + Preserve the declared activation rounding to the weight dtype while + using an identity VJP at that cast boundary for the FP64 golden graph. + """ + + hidden = validate_h3_adaln_projection(temb, weight, bias) + t64 = temb.double() + act = t64 * torch.sigmoid(t64) + # The declared cast rounds the value; its VJP is the identity. Written + # straight-through so autograd does not round the FP64 gradient to the + # weight dtype on its way back (``.to(bf16)`` would). + act = act + (act.to(weight.dtype).double() - act).detach() + table = F.linear(act, weight.double(), bias.double()) + return tuple(chunk.float() for chunk in split_adaln_table(table, hidden)) diff --git a/rl_engine/reference/minimax_h3/adaln_row_gather.py b/rl_engine/reference/minimax_h3/adaln_row_gather.py new file mode 100644 index 000000000..2cd713d45 --- /dev/null +++ b/rl_engine/reference/minimax_h3/adaln_row_gather.py @@ -0,0 +1,141 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 AdaLN row gather (RFC #420 row ``adaln_row_gather``). + +H3 computes one row index per packed-sequence position and uses it to pick +every position's six modulation vectors:: + + adaln_indices = timestep_indices * 3 + token_tags (S,) + shift_msa[adaln_indices], scale_msa[adaln_indices], ... six (S, H) + +``rows`` is the ``(3T, 6H)`` view of one block's AdaLN projection +(``table.view(-1, 6H)``); its six column blocks are the six modulation +tensors. ``timestep_indices`` and ``token_tags`` are semantic inputs +(RFC #420 §4): out-of-range values are an error, never clamped. +""" + +from __future__ import annotations + +import torch + +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM + +_INDEX_DTYPES = (torch.int64, torch.int32) + + +def validate_h3_adaln_row_gather( + rows: torch.Tensor, + timestep_indices: torch.Tensor, + token_tags: torch.Tensor, + *, + check_range: bool = True, +) -> int: + """Return the hidden size; raise on anything H3 does not define.""" + + if not isinstance(rows, torch.Tensor) or rows.dim() != 2: + raise ValueError("rows must be a 2-D (3T, 6H) tensor") + if not rows.is_floating_point(): + raise TypeError(f"rows must be floating point, got {rows.dtype}") + num_rows, width = rows.shape + if num_rows == 0 or num_rows % H3_MODALITY_NUM != 0: + raise ValueError(f"rows must hold 3 rows per timestep, got {num_rows}") + if width == 0 or width % H3_ADALN_CHUNKS != 0: + raise ValueError(f"rows width must be 6 * H, got {width}") + for name, index in (("timestep_indices", timestep_indices), ("token_tags", token_tags)): + if not isinstance(index, torch.Tensor) or index.dim() != 1: + raise ValueError(f"{name} must be a 1-D (S,) tensor") + if index.dtype not in _INDEX_DTYPES: + raise TypeError(f"{name} must be int64 or int32, got {index.dtype}") + if index.device != rows.device: + raise ValueError(f"{name} is on {index.device}, rows are on {rows.device}") + if timestep_indices.shape != token_tags.shape: + raise ValueError( + f"timestep_indices {tuple(timestep_indices.shape)} and token_tags " + f"{tuple(token_tags.shape)} must have the same length" + ) + if timestep_indices.numel() == 0: + raise ValueError("the packed sequence must not be empty") + if timestep_indices.dtype != token_tags.dtype: + raise TypeError("timestep_indices and token_tags must share an integer dtype") + if check_range: + num_timesteps = num_rows // H3_MODALITY_NUM + bad = ( + (timestep_indices < 0) + | (timestep_indices >= num_timesteps) + | (token_tags < 0) + | (token_tags >= H3_MODALITY_NUM) + ) + if bool(bad.any()): # one host sync + raise IndexError( + f"timestep_indices must lie in [0, {num_timesteps}) and token_tags in " + f"[0, {H3_MODALITY_NUM}) (0 video, 1 text, 2 audio)" + ) + return width // H3_ADALN_CHUNKS + + +def h3_adaln_indices(timestep_indices: torch.Tensor, token_tags: torch.Tensor) -> torch.Tensor: + return timestep_indices.long() * H3_MODALITY_NUM + token_tags.long() + + +class _DeterministicRowGather(torch.autograd.Function): + """``index_select`` forward; FP32 per-row segment sums for the backward. + + ``index_select``'s own backward is an atomic scatter-add in the grad dtype + (BF16 for H3), which is neither deterministic nor accurate. There are + only 3T table rows, so the reference sums each row's positions with one + FP32 ``torch.sum`` per row (a fixed reduction for a fixed packing) and + rounds once. + """ + + @staticmethod + def forward(ctx, rows, index): + ctx.save_for_backward(index) + ctx.rows_meta = (rows.shape[0], rows.dtype) + return torch.stack( + [chunk.index_select(0, index) for chunk in rows.chunk(H3_ADALN_CHUNKS, dim=-1)] + ) + + @staticmethod + def backward(ctx, grad): + (index,) = ctx.saved_tensors + num_rows, dtype = ctx.rows_meta + grad32 = grad.float() # (6, S, H) + out = grad32.new_zeros((num_rows, grad32.shape[0], grad32.shape[2])) + for row in range(num_rows): + positions = torch.nonzero(index == row).flatten() + if positions.numel(): + out[row] = grad32.index_select(1, positions).sum(dim=1) + return out.reshape(num_rows, -1).to(dtype), None + + +class NativeH3AdaLNRowGatherOp: + """PyTorch reference: ``index_select`` on each of the six column blocks. + + The forward is byte-identical to diffusers; the backward is deterministic + (see ``_DeterministicRowGather``). The raw diffusers path, including its + atomic backward, is ``rl_engine.validation.models.h3_provider.provider_adaln_row_gather``. + ``forward_fp32`` returns the same rows in FP32 with an FP64 gradient. + + The forward is a copy, but the VJP is a segmented reduction over the + packed sequence, so the op is judged with the ``reduction`` tolerance + class; forward bitwise equality is asserted separately. + """ + + op_class = "reduction" + + def __call__(self, rows, timestep_indices, token_tags): + return self.forward(rows, timestep_indices, token_tags) + + def forward(self, rows, timestep_indices, token_tags, *, check_range: bool = True): + validate_h3_adaln_row_gather(rows, timestep_indices, token_tags, check_range=check_range) + index = h3_adaln_indices(timestep_indices, token_tags) + return tuple(_DeterministicRowGather.apply(rows, index).unbind(0)) + + def forward_fp32(self, rows, timestep_indices, token_tags, *, check_range: bool = True): + validate_h3_adaln_row_gather(rows, timestep_indices, token_tags, check_range=check_range) + index = h3_adaln_indices(timestep_indices, token_tags) + return tuple( + chunk.index_select(0, index).float() + for chunk in rows.double().chunk(H3_ADALN_CHUNKS, dim=-1) + ) diff --git a/rl_engine/reference/minimax_h3/block.py b/rl_engine/reference/minimax_h3/block.py new file mode 100644 index 000000000..427d7de2e --- /dev/null +++ b/rl_engine/reference/minimax_h3/block.py @@ -0,0 +1,98 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""FP32 golden of one MiniMax-H3 transformer block (RFC #420 ``ws1_one_h3_block``). + +Every node takes FP32 inputs (BF16 weights and activations upcast exactly), +computes in FP32 with TF32 off, and returns FP32 without rounding back to +BF16, so the golden is the block's arithmetic without the BF16 storage the +provider and the candidate both have. It is the yardstick their errors are +measured against, not a bitwise target. The one declared cast kept from the +checkpoint contract is the AdaLN input: ``silu(temb)`` is rounded to BF16 +before the projection, as every implementation must do it. + +The partial MM-RoPE uses exact FP32 ``cos``/``sin`` tables on the leading 96 +of 128 head channels (the provider rounds the tables to BF16 first). +""" + +from __future__ import annotations + +import math + +import torch +import torch.nn.functional as F + +H3_HEADS = 56 +H3_HEAD_DIM = 128 +H3_ROTARY_DIM = 96 +H3_NORM_EPS = 1e-5 + + +def _f32(t: torch.Tensor) -> torch.Tensor: + return t.to(torch.float32) + + +def golden_adaln_tables(temb, weight, bias, hidden_size: int = 5376) -> tuple[torch.Tensor, ...]: + """Six ``(3T, H)`` FP32 tables; ``silu(temb)`` rounded to the weight dtype first.""" + + act = F.silu(_f32(temb)).to(weight.dtype) + table = F.linear(_f32(act), _f32(weight), _f32(bias)) + return table.view(-1, 6 * hidden_size).chunk(6, dim=-1) + + +def golden_rmsnorm(x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + x = _f32(x) + rstd = torch.rsqrt(x.square().mean(dim=-1, keepdim=True) + eps) + return x * rstd * _f32(weight) + + +def golden_norm_modulate(x, weight, shift, scale, index, eps: float = H3_NORM_EPS): + n = golden_rmsnorm(x, weight, eps) + return n * (1.0 + _f32(scale).index_select(0, index)) + _f32(shift).index_select(0, index) + + +def golden_linear(x, weight) -> torch.Tensor: + return F.linear(_f32(x), _f32(weight)) + + +def golden_rope_tables(position_ids, rope_freq_dim: int = 16, rope_theta: float = 10000.0): + """FP32 ``(S, 96)`` cos/sin, the same angles as the provider's rotary module.""" + + exponent = torch.arange( + 0, 2 * rope_freq_dim, 2, dtype=torch.float32, device=position_ids.device + ) + inv_freq = 1.0 / (rope_theta ** (exponent / (2 * rope_freq_dim))) + freqs = position_ids.to(torch.float32).unsqueeze(-1) * inv_freq.view(1, 1, -1) + freqs = torch.cat(freqs.unbind(dim=1), dim=-1) + freqs = torch.cat((freqs, freqs), dim=-1) + return freqs.cos(), freqs.sin() + + +def golden_apply_rotary(x, cos, sin) -> torch.Tensor: + """Rotate-half on the leading ``cos.shape[-1]`` channels of ``(B, S, H, D)``; FP32 tables.""" + + x = _f32(x) + rot = cos.shape[-1] + head, tail = x[..., :rot], x[..., rot:] + x1, x2 = head.chunk(2, dim=-1) + rotated = torch.cat((-x2, x1), dim=-1) + head = head * cos[None, :, None, :] + rotated * sin[None, :, None, :] + return torch.cat((head, tail), dim=-1) + + +def golden_attention(q, k, v) -> torch.Tensor: + """Non-causal full attention over one packed document, ``(B, S, H, D)`` -> ``(B, S, H*D)``.""" + + q, k, v = (_f32(t).permute(0, 2, 1, 3) for t in (q, k, v)) + scores = torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(q.shape[-1]) + out = torch.matmul(torch.softmax(scores, dim=-1), v) + return out.permute(0, 2, 1, 3).flatten(2, 3) + + +def golden_swiglu(projected) -> torch.Tensor: + hidden, gate = _f32(projected).chunk(2, dim=-1) + return hidden * F.silu(gate) + + +def golden_gate_residual(residual, gate, index, y) -> torch.Tensor: + return _f32(residual) + _f32(gate).index_select(0, index) * _f32(y) diff --git a/rl_engine/reference/minimax_h3/final_adaln_out.py b/rl_engine/reference/minimax_h3/final_adaln_out.py new file mode 100644 index 000000000..0b19aaee3 --- /dev/null +++ b/rl_engine/reference/minimax_h3/final_adaln_out.py @@ -0,0 +1,104 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 final AdaLN output (RFC #420 row ``final_adaln_out``). + +``MiniMaxH3AdaLayerNormOut`` runs after the 50 blocks:: + + shift, scale = linear(silu(temb).to(bf16)).chunk(2) # (T, H) each, shift first + out = norm(x) * (1.0 + scale[timestep_indices]) + shift[timestep_indices] + +``linear`` is ``norm_out.linear`` (``2688 -> 2 * 5376``, BF16), ``norm`` is +``norm_out.norm`` (``nn.RMSNorm(5376, eps=1e-5)``), and the table is indexed +by ``timestep_indices`` (one row per distinct timestep), not by +``adaln_indices``. Diffusers then casts the result to the FP32 output heads' +dtype, an exact upcast left to the caller. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from rl_engine.reference.minimax_h3.rmsnorm import ( + H3_NORM_EPS, + NativeH3RMSNormOp, + validate_h3_rmsnorm, +) + + +def validate_h3_final_adaln_out( + x, norm_weight, temb, weight, bias, timestep_indices, eps, *, same_dtype: bool = True +) -> int: + hidden = x.shape[-1] + if same_dtype: + validate_h3_rmsnorm(x, norm_weight, eps) + if temb.dtype != torch.float32: + raise TypeError(f"temb must be float32 (the SiLU runs before the cast), got {temb.dtype}") + if temb.dim() != 2 or temb.shape[0] == 0: + raise ValueError(f"temb must be a non-empty (T, D) matrix, got {tuple(temb.shape)}") + if weight.shape != (2 * hidden, temb.shape[1]) or bias.shape != (2 * hidden,): + raise ValueError( + f"norm_out.linear must be ({2 * hidden}, {temb.shape[1]}) with a ({2 * hidden},) " + f"bias, got {tuple(weight.shape)} / {tuple(bias.shape)}" + ) + if bias.dtype != weight.dtype: + raise TypeError("norm_out.linear weight and bias must share a dtype") + for name, tensor in (("temb", temb), ("weight", weight), ("bias", bias)): + if tensor.device != x.device: + raise ValueError(f"{name} is on {tensor.device}, x is on {x.device}") + if timestep_indices.dtype != torch.int64 or timestep_indices.dim() != 1: + raise TypeError("timestep_indices must be a 1-D int64 tensor") + if x.dim() < 2 or x.shape[-2] != timestep_indices.shape[0]: + raise ValueError("x must be (..., S, H) with S = len(timestep_indices)") + if bool(((timestep_indices < 0) | (timestep_indices >= temb.shape[0])).any()): + raise IndexError(f"timestep_indices must lie in [0, {temb.shape[0]})") + return hidden + + +class NativeH3FinalAdaLNOutOp: + """PyTorch reference: ``forward`` replays diffusers' ``norm_out``; + ``forward_fp32`` is the FP64 golden, returned in FP32. It rounds only where + the model itself stores a value in its own dtype (the SiLU cast, the BF16 + table, the ``norm_out.norm`` output and ``1 + scale``), straight-through for + the gradient, and computes everything else in FP64.""" + + op_class = "reduction" + + def __call__(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + return self.forward(x, norm_weight, temb, weight, bias, timestep_indices, eps) + + def forward(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + validate_h3_final_adaln_out(x, norm_weight, temb, weight, bias, timestep_indices, eps) + shift, scale = F.linear(F.silu(temb).to(weight.dtype), weight, bias).chunk(2, dim=-1) + return NativeH3RMSNormOp().forward_modulated( + x, norm_weight, shift, scale, timestep_indices, eps + ) + + def forward_fp32(self, x, norm_weight, temb, weight, bias, timestep_indices, eps=H3_NORM_EPS): + validate_h3_final_adaln_out( + x, norm_weight, temb, weight, bias, timestep_indices, eps, same_dtype=False + ) + t64 = temb.double() + act = t64 * torch.sigmoid(t64) + act = act + (act.to(weight.dtype).double() - act).detach() # declared cast, identity VJP + table = F.linear(act, weight.double(), bias.double()) + # norm_out.linear is a BF16 module: its (T, 2H) output is rounded to the + # weight dtype, and every position of a timestep shares that row, so the + # rounding is model semantics (straight-through for the gradient). + table = table + (table.to(weight.dtype).double() - table).detach() + shift, scale = table.chunk(2, dim=-1) + x64 = x.double() + n = x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + eps) * norm_weight.double() + # norm_out.norm is an x.dtype module and diffusers forms 1 + scale in + # x.dtype, so both are stored rounded. 1 + scale is shared by every + # position of a timestep: left unrounded, its error is systematic and + # grows with S in d_norm_weight; n's rounding enters d_scale (and so + # dW, d_temb) summed over S. Identity in FP32. + n = n + (n.to(x.dtype).double() - n).detach() + one_plus_scale = 1.0 + scale.index_select(0, timestep_indices) + one_plus_scale = ( + one_plus_scale + (one_plus_scale.to(x.dtype).double() - one_plus_scale).detach() + ) + out = n * one_plus_scale + shift.index_select(0, timestep_indices) + return out.float() diff --git a/rl_engine/reference/minimax_h3/fixed_order.py b/rl_engine/reference/minimax_h3/fixed_order.py new file mode 100644 index 000000000..e95a66054 --- /dev/null +++ b/rl_engine/reference/minimax_h3/fixed_order.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Shape-independent FP32 reductions shared by the H3 goldens and backwards. + +``torch.sum`` picks its reduction tree from the tensor shape, so a row's +result can change with the number of rows next to it. These helpers spell +the order out with elementwise ops only, which keeps every row's bytes a +function of that row alone. +""" + +from __future__ import annotations + +import torch + + +def tree_sum_lastdim_fp32(x: torch.Tensor) -> torch.Tensor: + """Pairwise tree over the last dim: pad to a power of two, halve until 1. + + Level ``l`` adds element ``i`` and ``i + width/2`` for every ``i``. The + tree depends only on the last-dim size. + """ + + x = x.float() + width = x.shape[-1] + if width == 0: + return x.new_zeros(x.shape[:-1]) + padded = 1 << (width - 1).bit_length() + if padded != width: + x = torch.nn.functional.pad(x, (0, padded - width)) + while x.shape[-1] > 1: + half = x.shape[-1] // 2 + x = x[..., :half] + x[..., half:] + return x[..., 0] + + +def fold_rows_fp32(rows: torch.Tensor) -> torch.Tensor: + """Left fold over dim 0 in ascending row order: ((r0 + r1) + r2) + ...""" + + rows = rows.float() + if rows.shape[0] == 0: + return rows.new_zeros(rows.shape[1:]) + acc = rows[0].clone() + for index in range(1, rows.shape[0]): + acc = acc + rows[index] + return acc diff --git a/rl_engine/reference/minimax_h3/gate_residual.py b/rl_engine/reference/minimax_h3/gate_residual.py new file mode 100644 index 000000000..f68a0fdf4 --- /dev/null +++ b/rl_engine/reference/minimax_h3/gate_residual.py @@ -0,0 +1,97 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 gated residual (RFC #420 row ``adaln_gate_residual``). + +After each sublayer of a block, diffusers adds the sublayer output back to +the residual stream through the per-row AdaLN gate:: + + hidden = residual + gate.index_select(0, adaln_indices) * sublayer_output + +``gate`` is ``gate_msa`` (after attention) or ``gate_mlp`` (after the FFN), +an ``(R, H)`` row view of the AdaLN table; ``residual`` and the sublayer +output are ``(..., S, H)``. The order is fixed by RFC #420 §4: +``residual + gate[row] * sublayer_output``. +""" + +from __future__ import annotations + +import torch + +_DTYPES = (torch.float32, torch.bfloat16, torch.float16) + + +def validate_h3_gate_residual( + residual: torch.Tensor, + y: torch.Tensor, + gate: torch.Tensor, + index: torch.Tensor, + *, + check_range: bool = True, + same_dtype: bool = True, +) -> None: + if residual.dtype not in _DTYPES: + raise TypeError(f"residual must be fp32, bf16 or fp16, got {residual.dtype}") + if y.shape != residual.shape or y.dtype != residual.dtype or y.device != residual.device: + raise ValueError("y must match the residual's shape, dtype and device") + if residual.dim() < 2 or residual.numel() == 0: + raise ValueError( + f"residual must be a non-empty (..., S, H) tensor, got {tuple(residual.shape)}" + ) + hidden = residual.shape[-1] + if gate.dim() != 2 or gate.shape[1] != hidden or gate.device != residual.device: + raise ValueError(f"gate must be (R, {hidden}) on the residual's device") + if same_dtype and gate.dtype != residual.dtype: + raise TypeError(f"gate dtype {gate.dtype} must match the residual dtype {residual.dtype}") + if index.dtype != torch.int64 or index.dim() != 1 or index.device != residual.device: + raise TypeError("index must be a 1-D int64 tensor on the residual's device") + if residual.shape[-2] != index.shape[0]: + raise ValueError(f"residual must be (..., S, H) with S = len(index) = {index.shape[0]}") + if check_range and bool(((index < 0) | (index >= gate.shape[0])).any()): + raise IndexError(f"index must lie in [0, {gate.shape[0]})") + + +class _DeterministicGateResidual(torch.autograd.Function): + """Eager forward; the gate gradient as FP32 per-row sums instead of + ``index_select``'s BF16 atomic scatter-add (``d_residual``/``d_y`` are the + eager VJPs).""" + + @staticmethod + def forward(ctx, residual, y, gate, index): + gathered = gate.index_select(0, index) + ctx.save_for_backward(y, gathered, index) + ctx.gate_rows = gate.shape[0] + return residual + gathered * y + + @staticmethod + def backward(ctx, grad): + y, gathered, index = ctx.saved_tensors + contrib = (grad.float() * y.float()).reshape(-1, index.shape[0], y.shape[-1]) + per_position = contrib.sum(dim=0) # (S, H): fixed reduction over the batch + d_gate = per_position.new_zeros((ctx.gate_rows, y.shape[-1])) + for row in range(ctx.gate_rows): + positions = torch.nonzero(index == row).flatten() + if positions.numel(): + d_gate[row] = per_position.index_select(0, positions).sum(dim=0) + return grad, grad * gathered, d_gate.to(gathered.dtype), None + + +class NativeH3GateResidualOp: + """PyTorch reference: ``forward`` is the eager provider expression (with a + deterministic gate gradient), ``forward_fp32`` the FP64 golden returned in + FP32. The raw diffusers path, atomic backward included, is + ``rl_engine.validation.models.h3_provider.provider_gate_residual``.""" + + op_class = "reduction" # forward is elementwise; the gate VJP sums positions + + def __call__(self, residual, y, gate, index): + return self.forward(residual, y, gate, index) + + def forward(self, residual, y, gate, index, *, check_range: bool = True): + validate_h3_gate_residual(residual, y, gate, index, check_range=check_range) + return _DeterministicGateResidual.apply(residual, y, gate, index) + + def forward_fp32(self, residual, y, gate, index): + validate_h3_gate_residual(residual, y, gate, index, same_dtype=False) + out = residual.double() + gate.double().index_select(0, index) * y.double() + return out.float() diff --git a/rl_engine/reference/minimax_h3/rmsnorm.py b/rl_engine/reference/minimax_h3/rmsnorm.py new file mode 100644 index 000000000..1727fd1d0 --- /dev/null +++ b/rl_engine/reference/minimax_h3/rmsnorm.py @@ -0,0 +1,153 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 RMSNorm and its fused AdaLN modulation (RFC #420 row ``h3_rmsnorm``). + +Every H3 RMSNorm is ``nn.RMSNorm(5376, eps=1e-5)`` with an affine weight in the +block dtype: the block ``norm1``/``norm2``, the token-refiner norms and the +final ``norm_out.norm``. Inside a transformer block (and in ``norm_out``) the +normalised rows are immediately modulated by per-row AdaLN parameters:: + + n = rms_norm(x, weight, eps) + out = n * (1.0 + scale[index]) + shift[index] # eager, rounded after every op + +``index`` is ``adaln_indices`` in a block and ``timestep_indices`` in +``norm_out``; ``shift``/``scale`` are ``(R, H)`` rows (views of the AdaLN +table). ``x`` is ``(..., S, H)`` and ``index`` has one entry per position ``S``. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + +H3_NORM_EPS = 1e-5 +_DTYPES = (torch.float32, torch.bfloat16, torch.float16) + + +def validate_h3_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> int: + for name, tensor in (("x", x), ("weight", weight)): + if not isinstance(tensor, torch.Tensor): + raise TypeError(f"{name} must be a torch.Tensor") + if tensor.dtype not in _DTYPES: + raise TypeError(f"{name} must be fp32, bf16 or fp16, got {tensor.dtype}") + if weight.dtype != x.dtype: + raise TypeError(f"weight dtype {weight.dtype} must match x dtype {x.dtype}") + if weight.device != x.device: + raise ValueError(f"weight is on {weight.device}, x is on {x.device}") + if x.dim() < 2 or x.numel() == 0: + raise ValueError(f"x must be a non-empty (..., N) tensor, got {tuple(x.shape)}") + if weight.shape != (x.shape[-1],): + raise ValueError(f"weight must be ({x.shape[-1]},), got {tuple(weight.shape)}") + if not eps > 0: + raise ValueError(f"eps must be positive, got {eps}") + return x.shape[-1] + + +def validate_h3_modulation( + x: torch.Tensor, + shift: torch.Tensor, + scale: torch.Tensor, + index: torch.Tensor, + *, + check_range: bool = True, + same_dtype: bool = True, +) -> None: + hidden = x.shape[-1] + for name, tensor in (("shift", shift), ("scale", scale)): + if tensor.device != x.device or (same_dtype and tensor.dtype != x.dtype): + raise TypeError(f"{name} must match x's dtype and device") + if tensor.dim() != 2 or tensor.shape[1] != hidden: + raise ValueError(f"{name} must be (R, {hidden}), got {tuple(tensor.shape)}") + if shift.shape != scale.shape: + raise ValueError("shift and scale must have the same shape") + if index.dtype != torch.int64 or index.dim() != 1 or index.device != x.device: + raise TypeError("index must be a 1-D int64 tensor on x's device") + if x.dim() < 2 or x.shape[-2] != index.shape[0]: + raise ValueError( + f"x must be (..., S, {hidden}) with S = len(index) = {index.shape[0]}, " + f"got {tuple(x.shape)}" + ) + if check_range and bool(((index < 0) | (index >= shift.shape[0])).any()): + raise IndexError(f"index must lie in [0, {shift.shape[0]})") + + +class _DeterministicModulation(torch.autograd.Function): + """Eager ``n * (1.0 + scale[i]) + shift[i]``; the ``shift``/``scale`` + gradients as FP32 per-row sums instead of ``index_select``'s atomic + scatter-add (``d_n`` is the eager VJP).""" + + @staticmethod + def forward(ctx, n, shift, scale, index): + t1 = 1.0 + scale.index_select(0, index) + ctx.save_for_backward(n, t1, index) + ctx.rows = shift.shape[0] + return n * t1 + shift.index_select(0, index) + + @staticmethod + def backward(ctx, grad): + n, t1, index = ctx.saved_tensors + hidden = n.shape[-1] + g = grad.float().reshape(-1, index.shape[0], hidden) + per_shift = g.sum(dim=0) + per_scale = (g * n.float().reshape(-1, index.shape[0], hidden)).sum(dim=0) + d_shift = per_shift.new_zeros((ctx.rows, hidden)) + d_scale = per_scale.new_zeros((ctx.rows, hidden)) + for row in range(ctx.rows): + positions = torch.nonzero(index == row).flatten() + if positions.numel(): + d_shift[row] = per_shift.index_select(0, positions).sum(dim=0) + d_scale[row] = per_scale.index_select(0, positions).sum(dim=0) + return grad * t1, d_shift.to(t1.dtype), d_scale.to(t1.dtype), None + + +class NativeH3RMSNormOp: + """PyTorch reference. + + ``forward`` / ``forward_modulated`` are the provider path (``F.rms_norm`` and + the eager modulation expression, with a deterministic shift/scale gradient; + the raw diffusers path is ``h3_provider.provider_norm_modulate``). ``forward_fp32`` / + ``forward_modulated_fp32`` are the golden: the same math in FP64, returned in + FP32. The modulated golden rounds, straight-through, only the two values the + model stores in ``x``'s dtype: ``norm(x)`` and ``1 + scale``. + """ + + op_class = "reduction" + + def __call__(self, x, weight, eps: float = H3_NORM_EPS): + return self.forward(x, weight, eps) + + def forward(self, x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + hidden = validate_h3_rmsnorm(x, weight, eps) + return F.rms_norm(x, (hidden,), weight, eps) + + def forward_fp32(self, x, weight, eps: float = H3_NORM_EPS) -> torch.Tensor: + validate_h3_rmsnorm(x, weight, eps) + x64 = x.double() + rstd = torch.rsqrt(x64.square().mean(dim=-1, keepdim=True) + eps) + return (x64 * rstd * weight.double()).float() + + def forward_modulated( + self, x, weight, shift, scale, index, eps: float = H3_NORM_EPS, *, check_range=True + ): + validate_h3_modulation(x, shift, scale, index, check_range=check_range) + n = self.forward(x, weight, eps) + return _DeterministicModulation.apply(n, shift, scale, index) + + def forward_modulated_fp32(self, x, weight, shift, scale, index, eps: float = H3_NORM_EPS): + # The golden may be fed FP32 golden modulation rows for a BF16 x. + validate_h3_modulation(x, shift, scale, index, same_dtype=False) + validate_h3_rmsnorm(x, weight, eps) + x64 = x.double() + n = x64 * torch.rsqrt(x64.square().mean(dim=-1, keepdim=True) + eps) * weight.double() + # The norm is an x.dtype module and diffusers forms 1 + scale in x.dtype. + # 1 + scale is shared by every position of a row, so leaving it unrounded + # skews d_norm_weight systematically with S; n's rounding enters d_scale + # summed over S. Identity in FP32. + n = n + (n.to(x.dtype).double() - n).detach() + one_plus_scale = 1.0 + scale.double().index_select(0, index) + one_plus_scale = ( + one_plus_scale + (one_plus_scale.to(x.dtype).double() - one_plus_scale).detach() + ) + out = n * one_plus_scale + shift.double().index_select(0, index) + return out.float() diff --git a/rl_engine/reference/minimax_h3/timestep_mlp.py b/rl_engine/reference/minimax_h3/timestep_mlp.py new file mode 100644 index 000000000..51ef48db7 --- /dev/null +++ b/rl_engine/reference/minimax_h3/timestep_mlp.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 FP32 timestep MLP (RFC #420 row ``timestep_mlp_fp32``). + +``TimestepEmbedding(in_channels=256, time_embed_dim=5376, out_dim=2688)``:: + + temb = linear_2(silu(linear_1(features))) 256 -> 5376 -> 2688 + +``time_embedder`` is in ``_keep_in_fp32_modules``: weights, biases, +activations and the output are FP32, and ``temb`` stays FP32 because every +AdaLN projection applies its own SiLU to it before casting (RFC #420 §4). +A BF16 argument is a contract violation (RFC probe H7), not a lower-precision +variant, so it is rejected. +""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F + + +def validate_h3_timestep_mlp( + x: torch.Tensor, + w1: torch.Tensor, + b1: torch.Tensor, + w2: torch.Tensor, + b2: torch.Tensor, +) -> None: + """Check a same-device FP32 ``(T, K) -> (T, H) -> (T, D)`` MLP contract. + + Require ``T > 0``, weights of shapes ``(H, K)``/``(D, H)`` and biases + of shapes ``(H,)``/``(D,)``; reject lower-precision parameters or inputs. + """ + + named = {"x": x, "w1": w1, "b1": b1, "w2": w2, "b2": b2} + for name, tensor in named.items(): + if not isinstance(tensor, torch.Tensor): + raise TypeError(f"{name} must be a torch.Tensor") + if tensor.dtype != torch.float32: + raise TypeError( + f"{name} must be float32: the H3 time_embedder is an FP32 module " + f"(got {tensor.dtype})" + ) + if tensor.device != x.device: + raise ValueError(f"{name} is on {tensor.device}, x is on {x.device}") + if x.dim() != 2 or x.shape[0] == 0: + raise ValueError(f"x must be a non-empty (T, K) matrix, got {tuple(x.shape)}") + hidden, k_in = w1.shape if w1.dim() == 2 else (None, None) + if k_in != x.shape[1]: + raise ValueError(f"w1 must be (H, {x.shape[1]}), got {tuple(w1.shape)}") + if b1.shape != (hidden,): + raise ValueError(f"b1 must be ({hidden},), got {tuple(b1.shape)}") + if w2.dim() != 2 or w2.shape[1] != hidden: + raise ValueError(f"w2 must be (D, {hidden}), got {tuple(w2.shape)}") + if b2.shape != (w2.shape[0],): + raise ValueError(f"b2 must be ({w2.shape[0]},), got {tuple(b2.shape)}") + + +class NativeH3TimestepMLPOp: + """PyTorch reference for the H3 timestep MLP. + + ``forward`` is the provider path (``nn.Linear`` / ``F.silu`` in FP32; run + it with TF32 disabled, as the RL-Kernel contract requires). + ``forward_fp32`` is the independent golden: the same graph in FP64, + rounded once to FP32, so its own rounding error (~1e-16 relative) is far + below the FP32 contract and its reduction order does not matter. + """ + + op_class = "reduction" + + def __call__(self, x, w1, b1, w2, b2): + """Return FP32 ``(T, D)`` embeddings using the PyTorch provider graph.""" + + return self.forward(x, w1, b1, w2, b2) + + def forward(self, x, w1, b1, w2, b2) -> torch.Tensor: + """Validate and evaluate linear-SiLU-linear in FP32 on the input device. + + Inputs follow the ``(T, K) -> (T, H) -> (T, D)`` contract and retain + PyTorch autograd support; CUDA callers must disable TF32 for this path. + """ + + validate_h3_timestep_mlp(x, w1, b1, w2, b2) + return F.linear(F.silu(F.linear(x, w1, b1)), w2, b2) + + def forward_fp32(self, x, w1, b1, w2, b2) -> torch.Tensor: + """Evaluate the validated MLP in FP64 and round its ``(T, D)`` output to FP32.""" + + validate_h3_timestep_mlp(x, w1, b1, w2, b2) + z = F.linear(x.double(), w1.double(), b1.double()) + h = z * torch.sigmoid(z) + return F.linear(h, w2.double(), b2.double()).float() diff --git a/rl_engine/reference/minimax_h3/timestep_sinusoid.py b/rl_engine/reference/minimax_h3/timestep_sinusoid.py new file mode 100644 index 000000000..0d95e58f0 --- /dev/null +++ b/rl_engine/reference/minimax_h3/timestep_sinusoid.py @@ -0,0 +1,129 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""H3 sinusoidal timestep features (RFC #420 row ``timestep_sinusoid_h3``). + +H3 builds ``Timesteps(num_channels=256, flip_sin_to_cos=True, +downscale_freq_shift=0)`` and feeds it the *distinct* timestep values of the +packed sequence, unscaled in ``[0, 1]`` (``t = 1 - sigma``):: + + freq[k] = exp(-ln(10000) * k / 128) k = 0..127, FP32 + arg[t, k] = t * freq[k] FP32 + out[t] = [cos(arg[t]) | sin(arg[t])] (T, 256) FP32 + +The output is always FP32: ``time_embedder`` is a ``_keep_in_fp32_modules`` +module in the checkpoint, so the features never pass through BF16. +""" + +from __future__ import annotations + +import math + +import torch + +from rl_engine.reference.minimax_h3 import H3_FREQ_DIM, H3_MAX_PERIOD + +_TIMESTEP_DTYPES = (torch.float32, torch.bfloat16, torch.float16) + + +def validate_h3_timesteps( + timestep: torch.Tensor, num_channels: int, *, check_range: bool = True +) -> None: + """Fail closed on inputs the H3 convention does not define. + + ``check_range`` rejects values outside ``[0, 1]``: H3 consumes + ``t = 1 - sigma`` unscaled, so a ``t * 1000`` caller (RFC probe H10) is a + convention error, not a different valid input. It reads the tensor back + to the host, which is cheap for the handful of distinct timesteps. + """ + + if not isinstance(timestep, torch.Tensor): + raise TypeError("timestep must be a torch.Tensor") + if timestep.dim() != 1: + raise ValueError(f"timestep must be 1-D (num_timesteps,), got {tuple(timestep.shape)}") + if timestep.numel() == 0: + raise ValueError("timestep must hold at least one timestep") + if timestep.dtype not in _TIMESTEP_DTYPES: + raise TypeError(f"timestep must be fp32, bf16 or fp16, got {timestep.dtype}") + if num_channels <= 0 or num_channels % 2 != 0: + raise ValueError(f"num_channels must be a positive even number, got {num_channels}") + if check_range: + t32 = timestep.detach().float() + # One host sync: NaN fails both comparisons, so this also rejects it. + if not bool(((t32 >= 0) & (t32 <= 1)).all()): + raise ValueError( + "timestep must be finite and lie in [0, 1]: H3 consumes t = 1 - sigma " + f"unscaled (got {t32.cpu().tolist()[:8]})" + ) + + +class NativeH3TimestepSinusoidOp: + """PyTorch reference for the H3 sinusoidal timestep features. + + ``forward`` replays diffusers' ``get_timestep_embedding`` op for op on the + input's device (the provider path). ``forward_fp32`` is the independent + golden: the same formula evaluated in FP64 and rounded once to FP32. + """ + + op_class = "elementwise" + + def __call__(self, timestep: torch.Tensor, *, num_channels: int = H3_FREQ_DIM): + """Return FP32 sinusoidal features for nonempty timesteps in ``[0, 1]``.""" + + return self.forward(timestep, num_channels=num_channels) + + def forward( + self, + timestep: torch.Tensor, + *, + num_channels: int = H3_FREQ_DIM, + check_range: bool = True, + ) -> torch.Tensor: + """Return FP32 ``(T, num_channels)`` cosine-then-sine features on the input device. + + Require nonempty 1-D FP32, BF16 or FP16 timesteps and positive even + channels. Validate finite values in ``[0, 1]`` when ``check_range`` is + true and preserve PyTorch autograd through the provider formula. + """ + + validate_h3_timesteps(timestep, num_channels, check_range=check_range) + half = num_channels // 2 + # Same op sequence and dtypes as diffusers get_timestep_embedding with + # flip_sin_to_cos=True, downscale_freq_shift=0, scale=1. + exponent = -math.log(H3_MAX_PERIOD) * torch.arange( + start=0, end=half, dtype=torch.float32, device=timestep.device + ) + exponent = exponent / half + freq = torch.exp(exponent) + arg = timestep[:, None].float() * freq[None, :] + return torch.cat([torch.cos(arg), torch.sin(arg)], dim=-1) + + def forward_fp32( + self, + timestep: torch.Tensor, + *, + num_channels: int = H3_FREQ_DIM, + check_range: bool = True, + ) -> torch.Tensor: + """Evaluate the validated sinusoid formula in FP64 and round once to FP32. + + Return ``(T, num_channels)`` on the timestep device with cosine channels + followed by sine channels and gradients through the FP64 golden graph. + """ + + validate_h3_timesteps(timestep, num_channels, check_range=check_range) + half = num_channels // 2 + k = torch.arange(half, dtype=torch.float64, device=timestep.device) + freq = torch.exp(-math.log(H3_MAX_PERIOD) * k / half) + arg = timestep.double()[:, None] * freq[None, :] + return torch.cat([torch.cos(arg), torch.sin(arg)], dim=-1).float() + + @staticmethod + def frequencies_fp32(num_channels: int, device: torch.device | str) -> torch.Tensor: + """The FP32 frequency table the provider path multiplies by.""" + + half = num_channels // 2 + exponent = -math.log(H3_MAX_PERIOD) * torch.arange( + start=0, end=half, dtype=torch.float32, device=device + ) + return torch.exp(exponent / half) diff --git a/rl_engine/runtime/registry.py b/rl_engine/runtime/registry.py index 7751d5901..5fe5b715b 100644 --- a/rl_engine/runtime/registry.py +++ b/rl_engine/runtime/registry.py @@ -200,6 +200,43 @@ class OpBackend(Enum, metaclass=_KernelEnumMeta): CUDA_SM90_LM_HEAD = "rl_engine.kernels.ops.cuda.linear.lm_head.SM90LMHeadOp" CUDA_SM90_EMBEDDING = "rl_engine.kernels.ops.cuda.linear.embedding.SM90EmbeddingOp" + # MiniMax-H3 (RFC #420) conditioning path + PYTORCH_H3_TIMESTEP_SINUSOID = ( + "rl_engine.reference.minimax_h3.timestep_sinusoid.NativeH3TimestepSinusoidOp" + ) + CUDA_H3_TIMESTEP_SINUSOID = ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "timestep_sinusoid.H3TimestepSinusoidCudaOp" + ) + PYTORCH_H3_TIMESTEP_MLP = "rl_engine.reference.minimax_h3.timestep_mlp.NativeH3TimestepMLPOp" + CUDA_H3_TIMESTEP_MLP = ( + "rl_engine.backends.cuda.model_specific.minimax_h3.timestep_mlp.H3TimestepMLPCudaOp" + ) + PYTORCH_H3_ADALN_PROJECTION = ( + "rl_engine.reference.minimax_h3.adaln_projection.NativeH3AdaLNProjectionOp" + ) + CUDA_H3_ADALN_PROJECTION = ( + "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection.H3AdaLNProjectionCudaOp" + ) + PYTORCH_H3_ADALN_ROW_GATHER = ( + "rl_engine.reference.minimax_h3.adaln_row_gather.NativeH3AdaLNRowGatherOp" + ) + CUDA_H3_ADALN_ROW_GATHER = ( + "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather.H3AdaLNRowGatherCudaOp" + ) + PYTORCH_H3_RMSNORM = "rl_engine.reference.minimax_h3.rmsnorm.NativeH3RMSNormOp" + CUDA_H3_RMSNORM = "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp" + PYTORCH_H3_GATE_RESIDUAL = "rl_engine.reference.minimax_h3.gate_residual.NativeH3GateResidualOp" + CUDA_H3_GATE_RESIDUAL = ( + "rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual.H3GateResidualCudaOp" + ) + PYTORCH_H3_FINAL_ADALN_OUT = ( + "rl_engine.reference.minimax_h3.final_adaln_out.NativeH3FinalAdaLNOutOp" + ) + CUDA_H3_FINAL_ADALN_OUT = ( + "rl_engine.backends.cuda.model_specific.minimax_h3.final_adaln_out.H3FinalAdaLNOutCudaOp" + ) + def _default_semantic_descriptors() -> tuple[OperatorBackendDescriptor, ...]: return ( @@ -440,6 +477,7 @@ class KernelRegistry: """ def __init__(self): + """Initialize backend caches, capability contracts, and device-specific priorities.""" self._instance_cache: Dict[str, Any] = {} self._failed_backends: Set[str] = set() self.semantic = SemanticOperatorCatalog(_default_semantic_descriptors()) @@ -609,6 +647,31 @@ def __init__(self): OpBackend.TRITON_ROPE, OpBackend.PYTORCH_NATIVE_ROPE, ], + "timestep_sinusoid_h3": [ + OpBackend.CUDA_H3_TIMESTEP_SINUSOID, + OpBackend.PYTORCH_H3_TIMESTEP_SINUSOID, + ], + "timestep_mlp_fp32": [ + OpBackend.CUDA_H3_TIMESTEP_MLP, + OpBackend.PYTORCH_H3_TIMESTEP_MLP, + ], + "adaln_projection_3mod": [ + OpBackend.CUDA_H3_ADALN_PROJECTION, + OpBackend.PYTORCH_H3_ADALN_PROJECTION, + ], + "adaln_row_gather": [ + OpBackend.CUDA_H3_ADALN_ROW_GATHER, + OpBackend.PYTORCH_H3_ADALN_ROW_GATHER, + ], + "h3_rmsnorm": [OpBackend.CUDA_H3_RMSNORM, OpBackend.PYTORCH_H3_RMSNORM], + "adaln_gate_residual": [ + OpBackend.CUDA_H3_GATE_RESIDUAL, + OpBackend.PYTORCH_H3_GATE_RESIDUAL, + ], + "final_adaln_out": [ + OpBackend.CUDA_H3_FINAL_ADALN_OUT, + OpBackend.PYTORCH_H3_FINAL_ADALN_OUT, + ], }, "rocm": { "logp": [ @@ -649,6 +712,13 @@ def __init__(self): "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.TRITON_SILU, OpBackend.PYTORCH_NATIVE_SILU], "swiglu": [OpBackend.TRITON_SWIGLU, OpBackend.PYTORCH_NATIVE_SWIGLU], + "timestep_sinusoid_h3": [OpBackend.PYTORCH_H3_TIMESTEP_SINUSOID], + "timestep_mlp_fp32": [OpBackend.PYTORCH_H3_TIMESTEP_MLP], + "adaln_projection_3mod": [OpBackend.PYTORCH_H3_ADALN_PROJECTION], + "adaln_row_gather": [OpBackend.PYTORCH_H3_ADALN_ROW_GATHER], + "h3_rmsnorm": [OpBackend.PYTORCH_H3_RMSNORM], + "adaln_gate_residual": [OpBackend.PYTORCH_H3_GATE_RESIDUAL], + "final_adaln_out": [OpBackend.PYTORCH_H3_FINAL_ADALN_OUT], }, "musa": { "logp": [ @@ -697,6 +767,13 @@ def __init__(self): ], "silu": [OpBackend.TRITON_SILU, OpBackend.PYTORCH_NATIVE_SILU], "swiglu": [OpBackend.TRITON_SWIGLU, OpBackend.PYTORCH_NATIVE_SWIGLU], + "timestep_sinusoid_h3": [OpBackend.PYTORCH_H3_TIMESTEP_SINUSOID], + "timestep_mlp_fp32": [OpBackend.PYTORCH_H3_TIMESTEP_MLP], + "adaln_projection_3mod": [OpBackend.PYTORCH_H3_ADALN_PROJECTION], + "adaln_row_gather": [OpBackend.PYTORCH_H3_ADALN_ROW_GATHER], + "h3_rmsnorm": [OpBackend.PYTORCH_H3_RMSNORM], + "adaln_gate_residual": [OpBackend.PYTORCH_H3_GATE_RESIDUAL], + "final_adaln_out": [OpBackend.PYTORCH_H3_FINAL_ADALN_OUT], }, "cpu": { "logp": [OpBackend.PYTORCH_NATIVE], @@ -722,6 +799,13 @@ def __init__(self): "embedding": [OpBackend.PYTORCH_NATIVE_EMBEDDING], "silu": [OpBackend.PYTORCH_NATIVE_SILU], "swiglu": [OpBackend.PYTORCH_NATIVE_SWIGLU], + "timestep_sinusoid_h3": [OpBackend.PYTORCH_H3_TIMESTEP_SINUSOID], + "timestep_mlp_fp32": [OpBackend.PYTORCH_H3_TIMESTEP_MLP], + "adaln_projection_3mod": [OpBackend.PYTORCH_H3_ADALN_PROJECTION], + "adaln_row_gather": [OpBackend.PYTORCH_H3_ADALN_ROW_GATHER], + "h3_rmsnorm": [OpBackend.PYTORCH_H3_RMSNORM], + "adaln_gate_residual": [OpBackend.PYTORCH_H3_GATE_RESIDUAL], + "final_adaln_out": [OpBackend.PYTORCH_H3_FINAL_ADALN_OUT], }, # Ascend NPU: op types without an entry fall back to their CPU # candidates (see the runtime override below), so only diff --git a/rl_engine/validation/models/h3_block.py b/rl_engine/validation/models/h3_block.py new file mode 100644 index 000000000..c974d9aab --- /dev/null +++ b/rl_engine/validation/models/h3_block.py @@ -0,0 +1,685 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""One real MiniMax-H3 transformer block, node by node (RFC #420 ``ws1_one_h3_block``). + +The block is a fixed graph of 17 nodes (``NODES``), from the AdaLN projection +through attention and the FFN to the second gated residual, on block 0's +pinned weights and a packed FL2VA layout. Every node runs three ways: + +* provider: the diffusers op sequence (``h3_provider``), fed provider inputs; +* candidate: the RL-Kernel binding of the node, fed candidate inputs + (chained) and, separately, provider inputs (isolated); +* golden: the FP32 reference (``rl_engine.reference.minimax_h3.block``). + +Each node also names the RFC row that owns its arithmetic and how it is bound +today (``Binding.status``): + +* ``row`` -- that row's operator, in this PR stack; +* ``interim`` -- an existing deterministic RL-Kernel operator standing in + until the owning row lands (its arithmetic, speed and rounding are that + operator's, not the row's); +* ``reference`` -- no kernel exists yet; the node runs the provider replay + itself, so it is trivially bitwise and contributes no RL-Kernel evidence. + +The report gives, per node: repeat and batch-row bitwise checks, chained and +isolated agreement with the provider, error against the golden for both the +candidate and the provider, and the backend that ran. ``first_*_drift`` name +the first node (graph order) where a comparison stops being bitwise; the +backward report does the same for the gradients reaching each node's output, +in reverse graph order. +""" + +from __future__ import annotations + +import hashlib +import time +from contextlib import contextmanager +from dataclasses import dataclass +from typing import Any, Callable + +import torch + +from rl_engine.reference.minimax_h3 import block as gold +from rl_engine.runtime.registry import KernelRegistry +from rl_engine.validation.models import h3_provider as prov +from rl_engine.validation.models.h3_cases import H3_BLOCK_LAYOUTS, h3_fl2va_layout, h3_timesteps + +HIDDEN = 5376 +HEADS = 56 +HEAD_DIM = 128 +BLOCK = "transformer_blocks.0" +NUM_TIMESTEPS = 2 + + +@dataclass(frozen=True) +class Binding: + rfc_row: str + status: str # "row" | "interim" | "reference" + backend: str + + +@dataclass(frozen=True) +class Node: + name: str + inputs: tuple[str, ...] + binding: Binding + # (ops, ctx, *inputs) -> output + candidate: Callable[..., Any] + # (ctx, *inputs) -> output + provider: Callable[..., Any] + golden: Callable[..., Any] + # Promise checked by the tests: the isolated candidate equals the provider bitwise. + provider_bitwise_isolated: bool + + +class CandidateOps: + """The operators bound to the block's nodes, resolved once per run. + + The three rows of this stack are resolved through the registry, so the + report records what dispatch actually returned; the interim operators are + named directly because their registry entries (``det_gemm``, + ``attention``, ``swiglu``) prefer backends that are not available or not + deterministic on every GPU this replay runs on. + """ + + def __init__(self, registry: KernelRegistry | None = None) -> None: + from rl_engine.backends.cuda.activation.swiglu import SwiGLUCudaOp + from rl_engine.backends.cuda.attention.deterministic_attn import DeterministicAttentionOp + from rl_engine.backends.shared.triton.gemm.det_gemm import TritonDetGemmOp + + registry = registry or KernelRegistry() + self.projection = registry.get_op("adaln_projection_3mod", device="cuda") + self.norm = registry.get_op("h3_rmsnorm", device="cuda") + self.gate = registry.get_op("adaln_gate_residual", device="cuda") + self.gemm = TritonDetGemmOp() + self.attention = DeterministicAttentionOp() + self.swiglu = SwiGLUCudaOp() + + def backends(self) -> dict[str, str]: + return { + name: f"{type(op).__module__}.{type(op).__name__}" for name, op in vars(self).items() + } + + +def _w(ctx: dict[str, Any], name: str) -> torch.Tensor: + return ctx["weights"][f"{BLOCK}.{name}"] + + +def _heads(x: torch.Tensor) -> torch.Tensor: + return x.unflatten(-1, (HEADS, HEAD_DIM)) + + +class _InterimLinear(torch.autograd.Function): + """``y = x @ W^T`` on the Triton FP32-accumulate kernel, rounded to BF16 once. + + Interim binding for the block's GEMM rows. ``TritonDetGemmOp.__call__`` + reduces K as a tree with BF16 nodes (about 3x cuBLAS's error at these + shapes); ``_triton_gemm_fp32`` runs one fixed-order FP32 K loop (no + split-K), so each output row depends only on its own input row. The + backward uses the same kernel for ``dX`` (reduction over the output + features, row-local) and ``dW`` (reduction over tokens). + """ + + @staticmethod + def forward(ctx, x2, weight): + from rl_engine.backends.shared.triton.gemm.det_gemm import _triton_gemm_fp32 + + ctx.save_for_backward(x2, weight) + return _triton_gemm_fp32(x2, weight.t()).to(x2.dtype) + + @staticmethod + def backward(ctx, grad): + from rl_engine.backends.shared.triton.gemm.det_gemm import _triton_gemm_fp32 + + x2, weight = ctx.saved_tensors + grad = grad.contiguous() + dx = _triton_gemm_fp32(grad, weight).to(x2.dtype) + dw = _triton_gemm_fp32(grad.t(), x2).to(weight.dtype) + return dx, dw + + +def _cand_linear(ops: CandidateOps, x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + flat = x.reshape(-1, x.shape[-1]) + return _InterimLinear.apply(flat, weight).view(*x.shape[:-1], weight.shape[0]) + + +def _cand_attention(ops: CandidateOps, q, k, v) -> torch.Tensor: + q, k, v = (t.transpose(1, 2).contiguous() for t in (q, k, _heads(v))) + out = ops.attention(q, k, v, causal=False) + return out.transpose(1, 2).flatten(2, 3) + + +def _cand_swiglu(ops: CandidateOps, projected: torch.Tensor) -> torch.Tensor: + hidden, gate = projected.chunk(2, dim=-1) + return ops.swiglu(gate.contiguous(), hidden.contiguous()) + + +def _linear_node(name: str, source: str, weight: str, row: str, status: str) -> Node: + return Node( + name=name, + inputs=(source,), + binding=Binding( + row, status, "triton _det_gemm_fp32_kernel (FP32 K loop, one BF16 rounding)" + ), + candidate=lambda ops, ctx, x: _cand_linear(ops, x, _w(ctx, weight)), + provider=lambda ctx, x: prov.provider_linear(x, _w(ctx, weight)), + golden=lambda ctx, x: gold.golden_linear(x, _w(ctx, weight)), + provider_bitwise_isolated=False, # tree reduction vs cuBLAS + ) + + +def _norm_modulate_node(name: str, source: str, weight: str, shift: int, scale: int) -> Node: + return Node( + name=name, + inputs=(source, "adaln_projection"), + binding=Binding("h3_rmsnorm", "row", "H3RMSNormCudaOp.forward_modulated"), + candidate=lambda ops, ctx, x, t: ops.norm.forward_modulated( + x, _w(ctx, weight), t[shift], t[scale], ctx["index"] + ), + provider=lambda ctx, x, t: prov.provider_norm_modulate( + x, _w(ctx, weight), t[shift], t[scale], ctx["index"] + ), + golden=lambda ctx, x, t: gold.golden_norm_modulate( + x, _w(ctx, weight), t[shift], t[scale], ctx["index"] + ), + provider_bitwise_isolated=True, + ) + + +def _head_norm_node(name: str, source: str, weight: str) -> Node: + return Node( + name=name, + inputs=(source,), + # The row is someone else's; until it lands, this stack's plain H3 RMSNorm + # (bitwise equal to nn.RMSNorm at any width) runs on the 128-wide heads. + binding=Binding("h3_qk_rmsnorm_d128", "interim", "H3RMSNormCudaOp.forward"), + candidate=lambda ops, ctx, x: ops.norm(_heads(x), _w(ctx, weight)), + provider=lambda ctx, x: prov.provider_head_rmsnorm(_heads(x), _w(ctx, weight)), + golden=lambda ctx, x: gold.golden_rmsnorm(_heads(x), _w(ctx, weight)), + provider_bitwise_isolated=True, + ) + + +def _rope_node(name: str, source: str) -> Node: + return Node( + name=name, + inputs=(source,), + binding=Binding("h3_mm_rope_3axis_partial", "reference", "provider_apply_rotary"), + candidate=lambda ops, ctx, x: prov.provider_apply_rotary(x, *ctx["rope"]), + provider=lambda ctx, x: prov.provider_apply_rotary(x, *ctx["rope"]), + golden=lambda ctx, x: gold.golden_apply_rotary(x, *ctx["rope_fp32"]), + provider_bitwise_isolated=True, + ) + + +def _gate_node(name: str, residual: str, sublayer: str, gate: int) -> Node: + return Node( + name=name, + inputs=(residual, sublayer, "adaln_projection"), + binding=Binding("adaln_gate_residual", "row", "H3GateResidualCudaOp"), + candidate=lambda ops, ctx, r, y, t: ops.gate(r, y, t[gate], ctx["index"]), + provider=lambda ctx, r, y, t: prov.provider_gate_residual(r, t[gate], ctx["index"], y), + golden=lambda ctx, r, y, t: gold.golden_gate_residual(r, t[gate], ctx["index"], y), + provider_bitwise_isolated=True, + ) + + +def _adaln_params(ctx): + return _w(ctx, "adaln_proj.linear.weight"), _w(ctx, "adaln_proj.linear.bias") + + +NODES: list[Node] = [ + Node( + name="adaln_projection", + inputs=("temb",), + binding=Binding("adaln_projection_3mod", "row", "H3AdaLNProjectionCudaOp"), + candidate=lambda ops, ctx, temb: ops.projection(temb, *_adaln_params(ctx)), + provider=lambda ctx, temb: prov.provider_adaln_modulation(temb, *_adaln_params(ctx)), + golden=lambda ctx, temb: gold.golden_adaln_tables(temb, *_adaln_params(ctx)), + provider_bitwise_isolated=False, # tensor-core tree vs cuBLAS + ), + _norm_modulate_node("norm1", "hidden", "norm1.weight", shift=0, scale=1), + _linear_node("q_proj", "norm1", "attn.to_q.weight", "h3_qkv_gemm", "interim"), + _linear_node("k_proj", "norm1", "attn.to_k.weight", "h3_qkv_gemm", "interim"), + _linear_node("v_proj", "norm1", "attn.to_v.weight", "h3_qkv_gemm", "interim"), + _head_norm_node("q_norm", "q_proj", "attn.norm_q.weight"), + _head_norm_node("k_norm", "k_proj", "attn.norm_k.weight"), + _rope_node("rope_q", "q_norm"), + _rope_node("rope_k", "k_norm"), + Node( + name="attention", + inputs=("rope_q", "rope_k", "v_proj"), + binding=Binding("h3_full_attention", "interim", "DeterministicAttentionOp(causal=False)"), + candidate=lambda ops, ctx, q, k, v: _cand_attention(ops, q, k, v), + provider=lambda ctx, q, k, v: prov.provider_attention(q, k, _heads(v)), + golden=lambda ctx, q, k, v: gold.golden_attention(q, k, _heads(v)), + provider_bitwise_isolated=False, + ), + _linear_node("o_proj", "attention", "attn.to_out.0.weight", "h3_attention_o_gemm", "interim"), + _gate_node("residual_attn", "hidden", "o_proj", gate=2), + _norm_modulate_node("norm2", "residual_attn", "norm2.weight", shift=3, scale=4), + _linear_node("ffn_gate_up", "norm2", "ff.net.0.proj.weight", "h3_ffn_gate_up_gemm", "interim"), + Node( + name="swiglu", + inputs=("ffn_gate_up",), + binding=Binding("h3_swiglu", "interim", "SwiGLUCudaOp"), + candidate=lambda ops, ctx, x: _cand_swiglu(ops, x), + provider=lambda ctx, x: prov.provider_swiglu(x), + golden=lambda ctx, x: gold.golden_swiglu(x), + provider_bitwise_isolated=False, # one FP32 rounding vs the eager two + ), + _linear_node("ffn_down", "swiglu", "ff.net.2.weight", "h3_ffn_down_gemm", "interim"), + _gate_node("residual_mlp", "residual_attn", "ffn_down", gate=5), +] +NODE_NAMES = tuple(node.name for node in NODES) +OUTPUT = NODES[-1].name + +# Gradient leaves of the backward replay: the block input, temb and every block parameter. +PARAM_LEAVES = ( + "adaln_proj.linear.weight", + "adaln_proj.linear.bias", + "norm1.weight", + "attn.to_q.weight", + "attn.to_k.weight", + "attn.to_v.weight", + "attn.norm_q.weight", + "attn.norm_k.weight", + "attn.to_out.0.weight", + "norm2.weight", + "ff.net.0.proj.weight", + "ff.net.2.weight", +) + + +# --- Inputs --------------------------------------------------------------------------------- + + +def make_block_context( + weights: dict[str, torch.Tensor], *, layout: str, seed: int = 0, batch: int = 1 +) -> dict[str, Any]: + """Block-0 weights, a packed FL2VA layout, the timestep embedding and a seeded input. + + ``temb`` comes from the provider's conditioning path (``h3_chain`` covers + those rows); the block input is a seeded BF16 stand-in for the packed + projections, which are other rows. Batch rows ``b > 0`` get their own + seeded inputs; row 0 is the same tensor at every ``batch``. + """ + + packed = h3_fl2va_layout(**H3_BLOCK_LAYOUTS[layout]) + seq_len = packed["position_ids"].shape[0] + timestep = h3_timesteps(NUM_TIMESTEPS, seed=seed) + w = weights + with torch.no_grad(): + features = prov.provider_time_proj(timestep) + temb = prov.provider_time_embedder( + features, + w["time_embedder.linear_1.weight"], + w["time_embedder.linear_1.bias"], + w["time_embedder.linear_2.weight"], + w["time_embedder.linear_2.bias"], + ) + rows = [ + torch.randn((1, seq_len, HIDDEN), generator=torch.Generator().manual_seed(seed + 2 + b)) + for b in range(batch) + ] + hidden = torch.cat(rows).to(torch.bfloat16).cuda() + return { + "layout": layout, + "seq_len": seq_len, + "batch": batch, + "weights": weights, + "temb": temb, + "hidden": hidden, + "index": packed["timestep_indices"] * 3 + packed["token_tags"], + "position_ids": packed["position_ids"], + "rope": prov.provider_rope(packed["position_ids"]), + "rope_fp32": gold.golden_rope_tables(packed["position_ids"]), + } + + +@contextmanager +def strict_fp32_matmul(): + """Golden matmuls in full FP32: no TF32, no reduced-precision reductions.""" + + saved = (torch.backends.cuda.matmul.allow_tf32, torch.backends.cudnn.allow_tf32) + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + try: + yield + finally: + torch.backends.cuda.matmul.allow_tf32, torch.backends.cudnn.allow_tf32 = saved + + +# --- Graph execution ------------------------------------------------------------------------ + + +def run_graph(mode: str, ctx: dict[str, Any], ops: CandidateOps | None = None) -> dict[str, Any]: + """Every node's output for one mode ("candidate", "provider" or "golden").""" + + values: dict[str, Any] = {"hidden": ctx["hidden"], "temb": ctx["temb"]} + with strict_fp32_matmul(): + for node in NODES: + args = [values[name] for name in node.inputs] + if mode == "candidate": + values[node.name] = node.candidate(ops, ctx, *args) + elif mode == "provider": + values[node.name] = node.provider(ctx, *args) + else: + values[node.name] = node.golden(ctx, *args) + return values + + +def run_isolated( + ctx: dict[str, Any], ops: CandidateOps, provider: dict[str, Any] +) -> dict[str, Any]: + """Each candidate node fed the provider's inputs to that node.""" + + return { + node.name: node.candidate(ops, ctx, *[provider[name] for name in node.inputs]) + for node in NODES + } + + +def _flat(value: Any) -> list[torch.Tensor]: + if isinstance(value, torch.Tensor): + return [value] + return [t for item in value for t in _flat(item)] + + +def _batch_row(value: Any, row: int, batch: int) -> Any: + """Row ``row`` of a batched activation; AdaLN tables have no batch axis.""" + + if isinstance(value, torch.Tensor): + return value[row : row + 1] if value.dim() >= 3 and value.shape[0] == batch else value + return tuple(_batch_row(v, row, batch) for v in value) + + +def bitwise(lhs: Any, rhs: Any) -> bool: + return all(torch.equal(a, b) for a, b in zip(_flat(lhs), _flat(rhs), strict=True)) + + +def max_abs(lhs: Any, rhs: Any) -> float: + return max( + float((a.float() - b.float()).abs().max()) if a.numel() else 0.0 + for a, b in zip(_flat(lhs), _flat(rhs), strict=True) + ) + + +def first_differing_row(lhs: Any, rhs: Any) -> int | None: + """First packed row (sequence position) where two node outputs differ, or None.""" + + for a, b in zip(_flat(lhs), _flat(rhs), strict=True): + if a.dim() < 3: + if not torch.equal(a, b): + return -1 # a table row, not a sequence position + continue + diff = (a != b).reshape(a.shape[0], a.shape[1], -1).any(dim=-1).any(dim=0) + rows = diff.nonzero() + if rows.numel(): + return int(rows[0]) + return None + + +def rel_err(value: Any, golden: Any) -> float: + """``max|value - golden| / max|golden|`` over a node's tensors.""" + + num = den = 0.0 + for a, g in zip(_flat(value), _flat(golden), strict=True): + g32 = g.float() + num = max(num, float((a.float() - g32).abs().max())) + den = max(den, float(g32.abs().max())) + return num / max(den, 1e-30) + + +def digest(value: Any) -> str: + """sha256 over a node's tensors' bytes, for comparing runs across processes or hosts.""" + + h = hashlib.sha256() + for t in _flat(value): + h.update(t.detach().contiguous().view(torch.uint8).cpu().numpy().tobytes()) + return h.hexdigest() + + +def _time_ms(fn: Callable[[], Any], repeats: int) -> float: + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(repeats): + start = time.perf_counter() + fn() + torch.cuda.synchronize() + samples.append((time.perf_counter() - start) * 1e3) + return sorted(samples)[len(samples) // 2] + + +# --- Forward report ------------------------------------------------------------------------- + + +def run_block_case( + weights: dict[str, torch.Tensor], + *, + layout: str, + seed: int = 0, + ops: CandidateOps | None = None, + timing_repeats: int = 0, +) -> dict[str, Any]: + """Forward replay of block 0: per-node agreement, invariance, accuracy and first drift.""" + + ops = ops or CandidateOps() + ctx = make_block_context(weights, layout=layout, seed=seed, batch=1) + ctx2 = make_block_context(weights, layout=layout, seed=seed, batch=2) + with torch.no_grad(): + provider = run_graph("provider", ctx) + golden = run_graph("golden", ctx) + chained = run_graph("candidate", ctx, ops) + repeat = run_graph("candidate", ctx, ops) + batched = run_graph("candidate", ctx2, ops) + isolated = run_isolated(ctx, ops, provider) + + report: dict[str, Any] = { + "layout": layout, + "seq_len": ctx["seq_len"], + "seed": seed, + "backends": ops.backends(), + "nodes": [], + } + firsts = {"first_drift": None, "first_isolated_drift": None} + firsts.update({"first_repeat_drift": None, "first_batch_drift": None}) + for node in NODES: + name = node.name + row0 = _batch_row(batched[name], 0, 2) + entry = { + "node": name, + "rfc_row": node.binding.rfc_row, + "status": node.binding.status, + "backend": node.binding.backend, + "provider_bitwise_isolated_promised": node.provider_bitwise_isolated, + "repeat_bitwise_equal": bitwise(repeat[name], chained[name]), + "batch_row_bitwise_equal": bitwise(row0, chained[name]), + "chained_vs_provider": { + "bitwise_equal": bitwise(chained[name], provider[name]), + "max_abs": max_abs(chained[name], provider[name]), + "first_differing_row": first_differing_row(chained[name], provider[name]), + }, + "isolated_vs_provider": { + "bitwise_equal": bitwise(isolated[name], provider[name]), + "max_abs": max_abs(isolated[name], provider[name]), + }, + "candidate_rel_err_vs_golden": rel_err(chained[name], golden[name]), + "provider_rel_err_vs_golden": rel_err(provider[name], golden[name]), + "candidate_sha256": digest(chained[name]), + } + checks = { + "first_drift": entry["chained_vs_provider"]["bitwise_equal"], + "first_isolated_drift": entry["isolated_vs_provider"]["bitwise_equal"], + "first_repeat_drift": entry["repeat_bitwise_equal"], + "first_batch_drift": entry["batch_row_bitwise_equal"], + } + for key, ok in checks.items(): + if firsts[key] is None and not ok: + firsts[key] = name + report["nodes"].append(entry) + report.update(firsts) + report["output"] = { + "candidate_rel_err_vs_golden": rel_err(chained[OUTPUT], golden[OUTPUT]), + "provider_rel_err_vs_golden": rel_err(provider[OUTPUT], golden[OUTPUT]), + "candidate_vs_provider_max_abs": max_abs(chained[OUTPUT], provider[OUTPUT]), + } + report["permutation_probe"] = permutation_probe(ctx, ops, chained[OUTPUT], seed) + if timing_repeats: + with torch.no_grad(): + report["timing_ms"] = { + "candidate_forward": _time_ms( + lambda: run_graph("candidate", ctx, ops), timing_repeats + ), + "provider_forward": _time_ms(lambda: run_graph("provider", ctx), timing_repeats), + } + return report + + +def permutation_probe( + ctx: dict[str, Any], ops: CandidateOps, base_output: torch.Tensor, seed: int +) -> dict[str, Any]: + """RFC #420 control H-C3, reported and not promised. + + Permute the packed rows together with every per-row tensor (positions, + AdaLN indices); the block output should be the same permutation of the + original. Row-local nodes keep each row's bytes, but attention sums over + keys in packed order, so a different key order may round differently. + """ + + perm = torch.randperm(ctx["seq_len"], generator=torch.Generator().manual_seed(seed + 7)).cuda() + permuted = { + **ctx, + "hidden": ctx["hidden"][:, perm], + "index": ctx["index"][perm], + "position_ids": ctx["position_ids"][perm], + "rope": tuple(t[perm] for t in ctx["rope"]), + "rope_fp32": tuple(t[perm] for t in ctx["rope_fp32"]), + } + with torch.no_grad(): + out = run_graph("candidate", permuted, ops)[OUTPUT] + expected = base_output[:, perm] + return { + "bitwise_equal_after_permutation": bool(torch.equal(out, expected)), + "max_abs": max_abs(out, expected), + "differing_rows": int((out != expected).reshape(out.shape[1], -1).any(dim=-1).sum()), + } + + +# --- Backward report ------------------------------------------------------------------------ + + +def _backward( + mode: str, ctx: dict[str, Any], ops: CandidateOps | None, upstream: torch.Tensor +) -> tuple[dict[str, torch.Tensor], dict[str, list[torch.Tensor]]]: + """Gradients of every leaf, and the gradient reaching each node's output, for one mode. + + The golden runs on FP32 copies of the leaves, so its gradients are not + rounded to the BF16 parameter dtype. + """ + + def leaf(t: torch.Tensor) -> torch.Tensor: + t = t.detach().float() if mode == "golden" else t.detach() + return t.clone().requires_grad_(True) + + leaves = {name: leaf(ctx["weights"][f"{BLOCK}.{name}"]) for name in PARAM_LEAVES} + leaves["hidden"] = leaf(ctx["hidden"]) + leaves["temb"] = leaf(ctx["temb"]) + weights = {**ctx["weights"], **{f"{BLOCK}.{k}": leaves[k] for k in PARAM_LEAVES}} + run_ctx = {**ctx, "weights": weights, "hidden": leaves["hidden"], "temb": leaves["temb"]} + values = run_graph(mode, run_ctx, ops) + for node in NODES: + for t in _flat(values[node.name]): + if t.requires_grad: + t.retain_grad() + out = values[OUTPUT] + torch.autograd.backward(out, upstream.to(out.dtype)) + node_grads = { + node.name: [ + t.grad if t.grad is not None else torch.zeros_like(t) for t in _flat(values[node.name]) + ] + for node in NODES + } + return {name: t.grad for name, t in leaves.items()}, node_grads + + +def run_block_backward( + weights: dict[str, torch.Tensor], + *, + layout: str, + seed: int = 0, + ops: CandidateOps | None = None, + timing_repeats: int = 0, +) -> dict[str, Any]: + """Backward replay of block 0. + + Per leaf (block input, ``temb`` and the 12 block parameters): repeat + bitwise, and the candidate's and provider's error against the FP32 + golden. Per node, in reverse graph order: whether the gradient reaching + the node's output repeats bitwise, and whether batch row 0 of a two-row + batch gets the same bytes as the batch of one (rows are separate packed + documents, so every activation gradient is row-local). Parameter + gradients sum over the batch and are not part of the batch check. + """ + + ops = ops or CandidateOps() + ctx = make_block_context(weights, layout=layout, seed=seed, batch=1) + ctx2 = make_block_context(weights, layout=layout, seed=seed, batch=2) + generator = torch.Generator().manual_seed(seed + 11) + up_row0 = torch.randn((1, ctx["seq_len"], HIDDEN), generator=generator) + up_row1 = torch.randn((1, ctx["seq_len"], HIDDEN), generator=generator) + upstream = up_row0.bfloat16().cuda() + upstream2 = torch.cat([up_row0, up_row1]).bfloat16().cuda() + + golden, _ = _backward("golden", ctx, None, upstream) + provider, _ = _backward("provider", ctx, None, upstream) + first, first_nodes = _backward("candidate", ctx, ops, upstream) + second, second_nodes = _backward("candidate", ctx, ops, upstream) + batched, batched_nodes = _backward("candidate", ctx2, ops, upstream2) + + report: dict[str, Any] = { + "layout": layout, + "seq_len": ctx["seq_len"], + "seed": seed, + "leaves": {}, + "nodes": [], + } + for name in (*PARAM_LEAVES, "hidden", "temb"): + report["leaves"][name] = { + "repeat_bitwise_equal": bool(torch.equal(first[name], second[name])), + "candidate_rel_err_vs_golden": rel_err(first[name], golden[name]), + "provider_rel_err_vs_golden": rel_err(provider[name], golden[name]), + } + report["leaves"]["hidden"]["batch_row_bitwise_equal"] = bool( + torch.equal(batched["hidden"][:1], first["hidden"]) + ) + first_repeat = first_batch = None + for node in reversed(NODES): + name = node.name + row0 = [_batch_row(g, 0, 2) for g in batched_nodes[name]] + entry = { + "node": name, + "grad_repeat_bitwise_equal": bitwise(first_nodes[name], second_nodes[name]), + "grad_batch_row_bitwise_equal": ( + bitwise(row0, first_nodes[name]) if node.name != "adaln_projection" else None + ), + } + if first_repeat is None and not entry["grad_repeat_bitwise_equal"]: + first_repeat = name + if first_batch is None and entry["grad_batch_row_bitwise_equal"] is False: + first_batch = name + report["nodes"].append(entry) + report["first_grad_repeat_drift"] = first_repeat + report["first_grad_batch_drift"] = first_batch + if timing_repeats: + report["timing_ms"] = { + "candidate_forward_backward": _time_ms( + lambda: _backward("candidate", ctx, ops, upstream), timing_repeats + ), + "provider_forward_backward": _time_ms( + lambda: _backward("provider", ctx, None, upstream), timing_repeats + ), + } + return report diff --git a/rl_engine/validation/models/h3_cases.py b/rl_engine/validation/models/h3_cases.py new file mode 100644 index 000000000..92455bae0 --- /dev/null +++ b/rl_engine/validation/models/h3_cases.py @@ -0,0 +1,188 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Deterministic MiniMax-H3 (RFC #420) test inputs shared by tests and scripts.""" + +from __future__ import annotations + +import torch + + +def h3_timesteps(num: int, *, seed: int = 0, device: str = "cuda") -> torch.Tensor: + """Return seeded FP32 timesteps of shape ``(num,)`` on ``device``. + + Requires positive ``num``; the first value is 0 and the last is 1 when + ``num > 1``, with all remaining values sampled uniformly in [0, 1]. + """ + + generator = torch.Generator(device="cpu").manual_seed(seed) + t = torch.rand(num, generator=generator) + t[0] = 0.0 + if num > 1: + t[-1] = 1.0 + return t.to(device) + + +def h3_packed_layout( + seq_len: int, num_timesteps: int, *, seed: int = 0, device: str = "cuda" +) -> tuple[torch.Tensor, torch.Tensor]: + """Return seeded int64 timestep indices and modality tags, each shape ``(seq_len,)``. + + Each row independently samples a timestep in ``[0, num_timesteps)`` and + a video/text/audio tag in ``[0, 3)``. The first three rows cover all tags + when available; both tensors are returned on ``device``. + """ + + generator = torch.Generator(device="cpu").manual_seed(seed) + token_tags = torch.randint(0, 3, (seq_len,), generator=generator) + timestep_indices = torch.randint(0, num_timesteps, (seq_len,), generator=generator) + token_tags[: min(3, seq_len)] = torch.arange(min(3, seq_len)) + return timestep_indices.to(device), token_tags.to(device) + + +def h3_block_layout( + seq_len: int, num_timesteps: int, *, seed: int = 0, device: str = "cuda" +) -> tuple[torch.Tensor, torch.Tensor]: + """(timestep_indices, token_tags) for block-structured packing, as H3 packs requests. + + Each timestep owns one contiguous run of text, video and audio blocks (about + 5% / 80% / 15% of its tokens), so table rows form contiguous position runs. + ``h3_packed_layout`` is the interleaved stress case. + """ + + generator = torch.Generator(device="cpu").manual_seed(seed) + cuts = torch.sort(torch.randperm(seq_len - 1, generator=generator)[: num_timesteps - 1] + 1) + edges = [0, *cuts.values.tolist(), seq_len] + timestep_indices = torch.empty(seq_len, dtype=torch.long) + token_tags = torch.empty(seq_len, dtype=torch.long) + for t, (lo, hi) in enumerate(zip(edges, edges[1:])): + text, audio = (hi - lo) // 20, (hi - lo) * 3 // 20 + timestep_indices[lo:hi] = t + token_tags[lo:hi] = 0 # video + token_tags[lo : lo + text] = 1 + token_tags[hi - audio : hi] = 2 + return timestep_indices.to(device), token_tags.to(device) + + +# diffusers@f53d552 modular_pipelines/minimax_h3/before_denoise.py constants. +_ROPE_FRAME_RESCALE = 5.0 / 3.0 +_ROPE_FRAMES_PER_LATENT = (1, 4, 4, 4, 4) +_ROPE_SPATIAL_SCALE = 32 +VIDEO_TAG, TEXT_TAG, AUDIO_TAG = 0, 1, 2 + + +def _spatial_grid(dim: int, patch: int, sqrt_area: float) -> torch.Tensor: + """One aspect-normalised spatial axis, as numpy's ``linspace(endpoint=False)`` builds it.""" + + import numpy as np + + ratio = dim / sqrt_area + left = (1.0 - ratio) / 2.0 + grid = np.linspace(left, left + ratio, dim // patch, endpoint=False) * _ROPE_SPATIAL_SCALE + return torch.from_numpy(grid).to(torch.float64) + + +def h3_fl2va_layout( + *, + num_text_tokens: int, + num_latent_frames: int, + latent_height: int, + latent_width: int, + num_audio_latents: int, + audio_channels: int = 2, + patch: tuple[int, int] = (2, 2), + device: str = "cuda", +) -> dict[str, torch.Tensor]: + """Packed text/audio/video rows of a text-only FL2VA request, laid out as diffusers packs them. + + Replays ``position_ids``/``token_tags`` of the pinned pipeline's FL2VA + layout step (no keyframe anchors): text rows on the time axis at their row + index, channel-major audio rows pinned to the width extremes, then the video + frames on the non-uniform ``5/3 * (1, 4, 4, 4, 4)`` clock. ``position_ids`` + is ``(S, 3)`` float64 as the pipeline makes it; the rotary module casts it. + ``timestep_indices`` puts text on table row 0 and the generated media on + row 1, a two-timestep stand-in for the conditioning rows. + """ + + import numpy as np + + patch_h, patch_w = patch + rows_per_frame = (latent_height // patch_h) * (latent_width // patch_w) + num_audio_rows = num_audio_latents * audio_channels + num_video_rows = num_latent_frames * rows_per_frame + seq_len = num_text_tokens + num_audio_rows + num_video_rows + audio_start = num_text_tokens + video_start = audio_start + num_audio_rows + + position_ids = torch.zeros(seq_len, 3, dtype=torch.float64) + position_ids[:num_text_tokens, 0] = torch.arange(num_text_tokens, dtype=torch.float64) + sqrt_area = np.sqrt(latent_height * latent_width) + height_grid = _spatial_grid(latent_height, patch_h, sqrt_area) + width_grid = _spatial_grid(latent_width, patch_w, sqrt_area) + grids = torch.meshgrid(height_grid, width_grid, indexing="ij") + frame_grid = torch.stack([grid.reshape(-1) for grid in grids], dim=-1) + + audio_time = float(num_text_tokens) + torch.arange(num_audio_latents, dtype=torch.float64) + position_ids[audio_start:video_start, 0] = audio_time.repeat(audio_channels) + position_ids[audio_start:video_start, 2] = torch.cat( + [ + torch.full((num_audio_latents,), float(width_grid[0]), dtype=torch.float64), + torch.full( + (num_audio_rows - num_audio_latents,), float(width_grid[-1]), dtype=torch.float64 + ), + ] + ) + spans = torch.tensor( + [ + _ROPE_FRAME_RESCALE * _ROPE_FRAMES_PER_LATENT[i % len(_ROPE_FRAMES_PER_LATENT)] + for i in range(num_latent_frames) + ], + dtype=torch.float64, + ) + frame_time = float(num_text_tokens) + torch.cat( + [torch.zeros(1, dtype=torch.float64), spans[:-1].cumsum(0)] + ) + video = torch.empty(num_latent_frames, rows_per_frame, 3, dtype=torch.float64) + video[:, :, 0] = frame_time[:, None] + video[:, :, 1:] = frame_grid[None] + position_ids[video_start:] = video.reshape(-1, 3) + + token_tags = torch.full((seq_len,), VIDEO_TAG, dtype=torch.long) + token_tags[:num_text_tokens] = TEXT_TAG + token_tags[audio_start:video_start] = AUDIO_TAG + timestep_indices = torch.ones(seq_len, dtype=torch.long) + timestep_indices[:num_text_tokens] = 0 + return { + "position_ids": position_ids.to(device), + "token_tags": token_tags.to(device), + "timestep_indices": timestep_indices.to(device), + } + + +# Named FL2VA layouts for the one-block replay: (text, latent frames, latent H, latent W, audio). +H3_BLOCK_LAYOUTS = { + # 32 + 2*20 + 2*(8*12) = 264 rows: the CI-sized case. + "tiny": dict( + num_text_tokens=32, + num_latent_frames=2, + latent_height=16, + latent_width=24, + num_audio_latents=20, + ), + # 77 + 2*40 + 3*(16*16) = 925 rows: a non-power-of-two, three-frame case. + "small": dict( + num_text_tokens=77, + num_latent_frames=3, + latent_height=32, + latent_width=32, + num_audio_latents=40, + ), + # 120 + 2*80 + 5*(18*32) = 3160 rows: a 16:9 five-frame case. + "medium": dict( + num_text_tokens=120, + num_latent_frames=5, + latent_height=36, + latent_width=64, + num_audio_latents=80, + ), +} diff --git a/rl_engine/validation/models/h3_chain.py b/rl_engine/validation/models/h3_chain.py new file mode 100644 index 000000000..72cbc760a --- /dev/null +++ b/rl_engine/validation/models/h3_chain.py @@ -0,0 +1,451 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Stage-by-stage replay of the MiniMax-H3 conditioning chain (RFC #420). + +Each stage runs three ways on the same device: + +* provider: the diffusers op sequence (``rl_engine.validation.models.h3_provider``); +* candidate: the dispatched RL-Kernel op, fed the previous candidate output + (chained) and, separately, the previous provider output (isolated); +* golden: the FP32/FP64 reference, fed the previous golden output. + +A stage also declares what it promises, so the end-to-end test and the +evidence JSON check the same thing: whether its isolated output is bitwise +equal to the provider, and its tolerance against the golden (rows of +``tolerance_contract.json``). ``first_drift`` is the first stage whose chained +output differs from the provider; ``first_isolated_drift`` the first stage +that differs even on provider inputs. +""" + +from __future__ import annotations + +import platform +import subprocess +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable + +import torch + +from rl_engine.runtime.registry import KernelRegistry +from rl_engine.validation.models.h3_cases import h3_packed_layout, h3_timesteps +from rl_engine.validation.models.h3_provider import ( + provider_adaln_modulation, + provider_adaln_row_gather, + provider_final_adaln_out, + provider_gate_residual, + provider_norm_modulate, + provider_time_embedder, + provider_time_proj, +) + +REPO_ROOT = Path(__file__).resolve().parents[3] +HIDDEN = 5376 + + +@dataclass(frozen=True) +class Stage: + name: str + op_type: str + # (op, ctx, upstream) -> outputs; ``upstream`` is the previous stage's output. + candidate: Callable[[Any, dict, Any], Any] + provider: Callable[[dict, Any], Any] + golden: Callable[[Any, dict, Any], Any] + # Promises checked by tests/models/minimax_h3/test_h3_conditioning_e2e.py. + provider_bitwise_isolated: bool + golden_atol: float + golden_rtol: float + + +def mlp_params(ctx: dict[str, Any]) -> list[torch.Tensor]: + """Return the context's MLP weights and biases in linear_1, linear_2 call order.""" + + weights = ctx["weights"] + return [weights[f"time_embedder.linear_{i}.{p}"] for i in (1, 2) for p in ("weight", "bias")] + + +def adaln_params(ctx: dict[str, Any]) -> tuple[torch.Tensor, torch.Tensor]: + """Return the context's block-0 AdaLN projection weight and bias.""" + + prefix = "transformer_blocks.0.adaln_proj.linear" + return ctx["weights"][f"{prefix}.weight"], ctx["weights"][f"{prefix}.bias"] + + +STAGES: list[Stage] = [ + Stage( + name="timestep_sinusoid_h3", + op_type="timestep_sinusoid_h3", + candidate=lambda op, ctx, _up: op(ctx["timestep"]), + provider=lambda ctx, _up: provider_time_proj(ctx["timestep"]), + golden=lambda op, ctx, _up: op.forward_fp32(ctx["timestep"]), + provider_bitwise_isolated=True, + golden_atol=1e-5, # elementwise / float32 + golden_rtol=1e-5, + ), + Stage( + name="timestep_mlp_fp32", + op_type="timestep_mlp_fp32", + candidate=lambda op, ctx, up: op(up, *mlp_params(ctx)), + provider=lambda ctx, up: provider_time_embedder(up, *mlp_params(ctx)), + golden=lambda op, ctx, up: op.forward_fp32(up, *mlp_params(ctx)), + provider_bitwise_isolated=False, # reduction: different tree from cuBLAS + golden_atol=1e-4, # reduction / float32 + golden_rtol=1e-4, + ), + Stage( + name="adaln_projection_3mod", + op_type="adaln_projection_3mod", + candidate=lambda op, ctx, up: op(up, *adaln_params(ctx)), + provider=lambda ctx, up: provider_adaln_modulation(up, *adaln_params(ctx)), + golden=lambda op, ctx, up: op.forward_fp32(up, *adaln_params(ctx)), + provider_bitwise_isolated=False, # reduction: tensor-core tree differs from cuBLAS + golden_atol=5e-2, # reduction / bfloat16 + golden_rtol=2e-2, + ), + Stage( + name="adaln_row_gather", + op_type="adaln_row_gather", + candidate=lambda op, ctx, up: op.gather_chunks( + up, ctx["timestep_indices"], ctx["token_tags"] + ), + provider=lambda ctx, up: provider_adaln_row_gather( + up, ctx["timestep_indices"], ctx["token_tags"] + ), + golden=lambda op, ctx, up: op.forward_fp32( + torch.cat(list(up), dim=1), ctx["timestep_indices"], ctx["token_tags"] + ), + provider_bitwise_isolated=True, # a copy + golden_atol=5e-2, # carries the projection's BF16 rounding (reduction / bfloat16) + golden_rtol=2e-2, + ), +] + + +def adaln_indices(ctx: dict[str, Any]) -> torch.Tensor: + return ctx["timestep_indices"] * 3 + ctx["token_tags"] + + +def _norm1_modulated(ctx: dict[str, Any], modulation, call): + shift_msa, scale_msa = modulation[0], modulation[1] + weight = ctx["weights"]["transformer_blocks.0.norm1.weight"] + return call(ctx["hidden"], weight, shift_msa, scale_msa, adaln_indices(ctx)) + + +STAGES.append( + Stage( + name="h3_rmsnorm", + op_type="h3_rmsnorm", + # Block norm1 followed by the MSA shift/scale of the projection stage. + candidate=lambda op, ctx, _up: _norm1_modulated( + ctx, ctx["history"]["adaln_projection_3mod"], op.forward_modulated + ), + provider=lambda ctx, _up: _norm1_modulated( + ctx, ctx["history"]["adaln_projection_3mod"], provider_norm_modulate + ), + golden=lambda op, ctx, _up: _norm1_modulated( + ctx, ctx["history"]["adaln_projection_3mod"], op.forward_modulated_fp32 + ), + provider_bitwise_isolated=True, # replays nn.RMSNorm's reduction order + golden_atol=5e-2, # reduction / bfloat16 + golden_rtol=2e-2, + ) +) + +STAGES.append( + Stage( + name="adaln_gate_residual", + op_type="adaln_gate_residual", + # residual + gate_msa[row] * attention_output (the stand-in sublayer output). + candidate=lambda op, ctx, _up: op( + ctx["hidden"], + ctx["sublayer"], + ctx["history"]["adaln_projection_3mod"][2], + adaln_indices(ctx), + ), + provider=lambda ctx, _up: provider_gate_residual( + ctx["hidden"], + ctx["history"]["adaln_projection_3mod"][2], + adaln_indices(ctx), + ctx["sublayer"], + ), + golden=lambda op, ctx, _up: op.forward_fp32( + ctx["hidden"], + ctx["sublayer"], + ctx["history"]["adaln_projection_3mod"][2], + adaln_indices(ctx), + ), + provider_bitwise_isolated=True, # eager rounding order, elementwise + golden_atol=5e-2, # carries the gate's BF16 rounding (reduction / bfloat16) + golden_rtol=2e-2, + ) +) + + +def _final(ctx: dict[str, Any], call): + weights = ctx["weights"] + history = ctx["history"] + return call( + history["adaln_gate_residual"], # the residual stream after the gated sublayer + weights["norm_out.norm.weight"], + history["timestep_mlp_fp32"], # temb + weights["norm_out.linear.weight"], + weights["norm_out.linear.bias"], + ctx["timestep_indices"], + ) + + +STAGES.append( + Stage( + name="final_adaln_out", + op_type="final_adaln_out", + candidate=lambda op, ctx, _up: _final(ctx, op), + provider=lambda ctx, _up: _final(ctx, provider_final_adaln_out), + golden=lambda op, ctx, _up: _final(ctx, op.forward_fp32), + provider_bitwise_isolated=False, # shift/scale projection: tensor-core tree vs cuBLAS + golden_atol=5e-2, # reduction / bfloat16 + golden_rtol=2e-2, + ) +) + +# The backward replay covers the conditioning chain up to the row gather. +CONDITIONING_STAGES = ( + "timestep_sinusoid_h3", + "timestep_mlp_fp32", + "adaln_projection_3mod", + "adaln_row_gather", +) + +# Parameters whose gradients the backward replay reports, in chain order. +GRAD_LEAVES = ( + "time_embedder.linear_1.weight", + "time_embedder.linear_1.bias", + "time_embedder.linear_2.weight", + "time_embedder.linear_2.bias", + "transformer_blocks.0.adaln_proj.linear.weight", + "transformer_blocks.0.adaln_proj.linear.bias", +) +BACKWARD_MODES = ("candidate", "candidate_fused", "provider") + + +def _flat(value: Any) -> list[torch.Tensor]: + """Collect tensors from nested sequences in their original order.""" + + if isinstance(value, torch.Tensor): + return [value] + return [tensor for item in value for tensor in _flat(item)] + + +def compare(lhs: Any, rhs: Any, atol: float = 0.0, rtol: float = 0.0) -> dict[str, Any]: + """Report tensor equality, maximum FP32 absolute error, and tolerance checks.""" + + lhs_t, rhs_t = _flat(lhs), _flat(rhs) + bitwise = all(torch.equal(a, b) for a, b in zip(lhs_t, rhs_t, strict=True)) + max_abs = 0.0 + within = True + for a, b in zip(lhs_t, rhs_t, strict=True): + diff = (a.float() - b.float()).abs() + if diff.numel(): + max_abs = max(max_abs, float(diff.max())) + within = within and bool((diff <= atol + rtol * b.float().abs()).all()) + return {"bitwise_equal": bitwise, "max_abs": max_abs, "within_tolerance": within} + + +def make_context(weights, *, num_timesteps: int, seq_len: int, seed: int) -> dict[str, Any]: + """Build seeded CUDA timesteps and packed row labels with the supplied weights.""" + + timestep_indices, token_tags = h3_packed_layout(seq_len, num_timesteps, seed=seed) + return { + "timestep": h3_timesteps(num_timesteps, seed=seed), + "timestep_indices": timestep_indices, + "token_tags": token_tags, + "weights": weights, + # Stand-in for the packed hidden states (the patch/text projections are + # other RFC rows): BF16 (1, S, H), seeded. + "hidden": torch.randn( + (1, seq_len, HIDDEN), generator=torch.Generator().manual_seed(seed + 2) + ) + .to(torch.bfloat16) + .cuda(), + # Stand-in for the attention output the gate multiplies (the attention + # rows belong to other contributors): BF16 (1, S, H), seeded. + "sublayer": ( + torch.randn((1, seq_len, HIDDEN), generator=torch.Generator().manual_seed(seed + 3)) + * 3.0 + ) + .to(torch.bfloat16) + .cuda(), + } + + +def golden_op(registry: KernelRegistry, op_type: str): + """Return the PyTorch reference at the end of the operator's CPU priority list.""" + + return registry._get_or_create_backend(registry._priority_map["cpu"][op_type][-1]) + + +def run_case( + registry: KernelRegistry, + weights, + *, + num_timesteps: int, + seq_len: int, + seed: int = 0, + stages: list[Stage] | None = None, +) -> dict[str, Any]: + """Replay CUDA candidate, provider, and golden paths and report each stage's drift. + + Supplied stages must include their prerequisites in execution order. + Omitting ``stages`` runs the full registered chain without gradients. + """ + + stages = STAGES if stages is None else stages + ctx = make_context(weights, num_timesteps=num_timesteps, seq_len=seq_len, seed=seed) + report: dict[str, Any] = { + "num_timesteps": num_timesteps, + "seq_len": seq_len, + "seed": seed, + "stages": [], + } + chained = provider = golden = None + first_drift = first_isolated = None + # Each mode sees its own earlier outputs (a stage may read any earlier stage). + history = {"candidate": {}, "provider": {}, "golden": {}} + ctx_c = {**ctx, "history": history["candidate"]} + ctx_p = {**ctx, "history": history["provider"]} + ctx_g = {**ctx, "history": history["golden"]} + for stage in stages: + op = registry.get_op(stage.op_type, device="cuda") + gold = golden_op(registry, stage.op_type) + with torch.no_grad(): + isolated = stage.candidate(op, ctx_p, provider) + chained_input = chained + chained = stage.candidate(op, ctx_c, chained_input) + repeat = stage.candidate(op, ctx_c, chained_input) + golden = stage.golden(gold, ctx_g, golden) + provider = stage.provider(ctx_p, provider) + history["candidate"][stage.name] = chained + history["provider"][stage.name] = provider + history["golden"][stage.name] = golden + entry = { + "stage": stage.name, + "backend": type(op).__name__, + "kernel_id": getattr(op, "kernel_id", type(op).__name__), + "repeat_bitwise_equal": compare(repeat, chained)["bitwise_equal"], + "chained_vs_provider": compare(chained, provider), + "isolated_vs_provider": compare(isolated, provider), + "chained_vs_golden": compare(chained, golden, stage.golden_atol, stage.golden_rtol), + "provider_vs_golden": compare(provider, golden, stage.golden_atol, stage.golden_rtol), + } + if first_drift is None and not entry["chained_vs_provider"]["bitwise_equal"]: + first_drift = stage.name + if first_isolated is None and not entry["isolated_vs_provider"]["bitwise_equal"]: + first_isolated = stage.name + report["stages"].append(entry) + report["first_drift"] = first_drift + report["first_isolated_drift"] = first_isolated + return report + + +def chain_grads(mode: str, registry: KernelRegistry, ctx, upstream) -> list[torch.Tensor]: + """Parameter gradients of the full chain for one execution mode. + + ``candidate`` runs the four RL-Kernel ops separately, ``candidate_fused`` + runs projection + gather as ``H3AdaLNModulationCudaOp``, ``provider`` the + diffusers path, ``golden`` the FP64 references. + """ + + leaves = { + name: ctx["weights"][name].detach().clone().requires_grad_(True) for name in GRAD_LEAVES + } + run_ctx = {**ctx, "weights": {**ctx["weights"], **leaves}, "history": {}} + value = None + if mode == "candidate_fused": + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_modulation import ( + H3AdaLNModulationCudaOp, + ) + + for stage in STAGES[:2]: + value = stage.candidate(registry.get_op(stage.op_type, device="cuda"), run_ctx, value) + value = H3AdaLNModulationCudaOp()( + value, *adaln_params(run_ctx), ctx["timestep_indices"], ctx["token_tags"] + ) + else: + for stage in (s for s in STAGES if s.name in CONDITIONING_STAGES): + if mode == "provider": + value = stage.provider(run_ctx, value) + elif mode == "candidate": + op = registry.get_op(stage.op_type, device="cuda") + value = stage.candidate(op, run_ctx, value) + else: + value = stage.golden(golden_op(registry, stage.op_type), run_ctx, value) + outputs = list(value) + torch.autograd.backward( + outputs, [g.to(out.dtype) for g, out in zip(upstream, outputs, strict=True)] + ) + return [leaves[name].grad for name in GRAD_LEAVES] + + +def run_backward_case( + registry: KernelRegistry, weights, *, num_timesteps: int, seq_len: int, seed: int = 0 +) -> dict[str, Any]: + """Determinism and accuracy of the chain's parameter gradients, per mode.""" + + ctx = make_context(weights, num_timesteps=num_timesteps, seq_len=seq_len, seed=seed) + generator = torch.Generator(device="cuda").manual_seed(seed + 1) + hidden = weights[GRAD_LEAVES[4]].shape[0] // 18 + upstream = [ + torch.randn(seq_len, hidden, device="cuda", generator=generator).bfloat16() + for _ in range(6) + ] + golden = chain_grads("golden", registry, ctx, upstream) + runs = { + mode: [chain_grads(mode, registry, ctx, upstream) for _ in range(2)] + for mode in BACKWARD_MODES + } + report: dict[str, Any] = {"num_timesteps": num_timesteps, "seq_len": seq_len, "leaves": {}} + for index, name in enumerate(GRAD_LEAVES): + gold = golden[index].float() + entry: dict[str, Any] = {"golden_absmax": float(gold.abs().max())} + for mode, (first, second) in runs.items(): + grad = first[index] + err = (grad.float() - gold).abs().max() + entry[mode] = { + "repeat_bitwise_equal": bool(torch.equal(first[index], second[index])), + "max_abs_vs_golden": float(err), + "max_abs_vs_golden_over_absmax": float(err / gold.abs().max().clamp_min(1e-30)), + "correctly_rounded_fraction": float( + (grad == golden[index].to(grad.dtype)).float().mean() + ), + } + report["leaves"][name] = entry + return report + + +def git_state() -> dict[str, Any]: + """Return HEAD and tracked-tree dirtiness, defaulting to unknown/dirty if Git fails.""" + + def run(*args: str) -> str: + """Run Git in the repository root and return stripped stdout.""" + + return subprocess.check_output(["git", *args], cwd=REPO_ROOT, text=True).strip() + + try: + sha = run("rev-parse", "HEAD") + dirty = bool(run("status", "--porcelain", "--untracked-files=no")) + except Exception: # noqa: BLE001 + sha, dirty = "unknown", True + return {"rl_kernel_commit": sha, "tracked_tree_dirty": dirty} + + +def environment() -> dict[str, Any]: + """Record the active CUDA device, software versions, and matmul TF32 setting.""" + + return { + "gpu": torch.cuda.get_device_name(), + "capability": list(torch.cuda.get_device_capability()), + "torch": torch.__version__, + "cuda": torch.version.cuda, + "python": platform.python_version(), + "tf32": bool(torch.backends.cuda.matmul.allow_tf32), + } diff --git a/rl_engine/validation/models/h3_manifest.json b/rl_engine/validation/models/h3_manifest.json new file mode 100644 index 000000000..e590a8473 --- /dev/null +++ b/rl_engine/validation/models/h3_manifest.json @@ -0,0 +1,137 @@ +{ + "version": "h3-conditioning-v1", + "rfc": "RL-Align/RL-Kernel#420", + "model_identity": { + "hf_repo": "MiniMaxAI/MiniMax-H3", + "revision": "42ed227ee7df40d41602854ae760620d6eb651fe", + "subfolder": "transformer", + "architecture": "MiniMaxH3Transformer3DModel", + "config_sha256": "74c11bff524336576096993cbfcdcdc2ef4fa2fa4409df693bdcbc6c666282ae", + "index_file": "diffusion_pytorch_model.safetensors.index.json", + "index_sha256": "ac30a3b58963f2e735d493475fbb81853a5735ec947619648b3e045acda6783e", + "config_fingerprint": { + "hidden_size": 5376, + "num_attention_heads": 56, + "attention_head_dim": 128, + "num_layers": 50, + "num_refiner_layers": 2, + "ffn_dim": 14336, + "freq_dim": 256, + "time_embed_hidden_dim": 5376, + "time_embed_dim": 2688, + "norm_eps": 1e-05, + "final_norm_eps": 1e-05 + } + }, + "weight_shards": { + "diffusion_pytorch_model-00001-of-00014.safetensors": { + "sha256": "2d847200c45c09dd7f973c1b096663068408ef851ee0b3711d059b6dc5dcd028", + "size_bytes": 4825958704 + } + }, + "tensors": { + "time_embedder.linear_1.weight": { + "dtype": "float32", "shape": [5376, 256], + "sha256": "4c50eabe3c434dc05422e6501cdf7baa5cecbb1a7706ae9599c7a5b0e34e2913" + }, + "time_embedder.linear_1.bias": { + "dtype": "float32", "shape": [5376], + "sha256": "64763a5e39cb078eb0d8c5be5d56ad36cac3dec37be93a87e106d82ea6c071b7" + }, + "time_embedder.linear_2.weight": { + "dtype": "float32", "shape": [2688, 5376], + "sha256": "697819dd9413e0fec3a50987a964f1fd5db568ca545267414238f452294a8154" + }, + "time_embedder.linear_2.bias": { + "dtype": "float32", "shape": [2688], + "sha256": "096a3c7539f0248698cc65376c8da260e7365bf4ccef847c09bb9a36782a743c" + }, + "transformer_blocks.0.adaln_proj.linear.weight": { + "dtype": "bfloat16", "shape": [96768, 2688], + "sha256": "a0000193b30ea36a1996f53fc0d6c4c37ce5fd4585284f0f9bae570d25c6b8cf" + }, + "transformer_blocks.0.adaln_proj.linear.bias": { + "dtype": "bfloat16", "shape": [96768], + "sha256": "d59f9de9e07e6b4e3b57009627e276a26971ce379555fc99dbb06d3a237c8fc0" + }, + "transformer_blocks.0.norm1.weight": { + "dtype": "bfloat16", "shape": [5376], + "sha256": "5e32e4d7379a9f9302755af395ed50b5b998352a221dc2c99204a85dd9f9b49f" + }, + "transformer_blocks.0.norm2.weight": { + "dtype": "bfloat16", "shape": [5376], + "sha256": "a2c513ec459414e211a3d37f77c1f556f11117f56516b22eb5912f8ff5f98872" + }, + "token_refiner.final_norm.weight": { + "dtype": "bfloat16", "shape": [5376], + "sha256": "cf10a56a11370a216a4ac3afe9e8b8a3133206bccd6011a92b9aa791a9196b93" + }, + "norm_out.norm.weight": { + "dtype": "bfloat16", "shape": [5376], + "sha256": "91ac17792929e0a84533b28caf71ac831ec3d666c07a473393639be20392f5ce" + }, + "norm_out.linear.weight": { + "dtype": "bfloat16", "shape": [10752, 2688], + "sha256": "38063ae4b492aacd320e71d4a38a32325c0d47b29f59978f0f4979ca080fe673" + }, + "norm_out.linear.bias": { + "dtype": "bfloat16", "shape": [10752], + "sha256": "3b12cdc7bbda772e7e31539ffa9a32f529dba7f4148273d2128bdc621fd27438" + } + }, + "block_tensors": { + "transformer_blocks.0.attn.to_q.weight": { + "dtype": "bfloat16", "shape": [7168, 5376], + "sha256": "196f9120ad018784bcefd4224f0e08fae88ba1adf079ecdca18410a414a6a2a4" + }, + "transformer_blocks.0.attn.to_k.weight": { + "dtype": "bfloat16", "shape": [7168, 5376], + "sha256": "8e9171f3b8cf894904d1ddbd96510739d7291d2992310de6a52e1d733ec895cd" + }, + "transformer_blocks.0.attn.to_v.weight": { + "dtype": "bfloat16", "shape": [7168, 5376], + "sha256": "275bbc993b1ba738831e04d86f290c9988108d96fe8c97c3b576a043c7ff3fe1" + }, + "transformer_blocks.0.attn.norm_q.weight": { + "dtype": "bfloat16", "shape": [128], + "sha256": "34a3f13557a2c0a9af260e48d431041e381779a0f61616cfe9800ec839693519" + }, + "transformer_blocks.0.attn.norm_k.weight": { + "dtype": "bfloat16", "shape": [128], + "sha256": "8991949ddea3bc2136a0fcf0b980fa63bd0df760207e3116f7ce96df4e6ca058" + }, + "transformer_blocks.0.attn.to_out.0.weight": { + "dtype": "bfloat16", "shape": [5376, 7168], + "sha256": "822cb426aaebe57d3d219b3c4f6855ca5c13821cb84ebe0a7474fe3a87869fd6" + }, + "transformer_blocks.0.ff.net.0.proj.weight": { + "dtype": "bfloat16", "shape": [28672, 5376], + "sha256": "49587044f325aae73f046c4cda9637052670227ec3c2c289359d9d1ebf29f95d" + }, + "transformer_blocks.0.ff.net.2.weight": { + "dtype": "bfloat16", "shape": [5376, 14336], + "sha256": "e19d4e1649e4b55b1931842e5cc4bb46892c2e635b9315b2996630a86284e88f" + } + }, + "reference_implementation": { + "repo": "huggingface/diffusers", + "commit": "f53d552036a0d1bd5570782a39cd40cfabf112bc", + "file": "src/diffusers/models/transformers/transformer_minimax_h3.py", + "file_sha256": "435873d9bc26b8770342961e18cb20324aeeaa4d045e85e0941638d4857d7db3", + "timestep_embedding_file": "src/diffusers/models/embeddings.py" + }, + "conventions": { + "timestep": "distinct timesteps t = 1 - sigma, unscaled in [0, 1], shape (T,)", + "modality_tags": {"video": 0, "text": 1, "audio": 2}, + "adaln_row": "timestep_index * 3 + token_tag", + "adaln_table_rows": "[t0_mod0, t0_mod1, t0_mod2, t1_mod0, ...]", + "adaln_chunk_order": ["shift_msa", "scale_msa", "gate_msa", "shift_mlp", "scale_mlp", "gate_mlp"], + "precision": { + "timestep_sinusoid": "float32", + "timestep_mlp": "float32 weights, float32 activations", + "adaln_projection": "SiLU in float32, cast to bfloat16, bfloat16 weight/bias/output", + "rmsnorm": "nn.RMSNorm(5376, eps=1e-5), bfloat16 affine weight, statistics in float32", + "final_modulation": "norm_out.linear(silu(temb).bfloat16) -> (shift, scale), indexed by timestep_indices" + } + } +} diff --git a/rl_engine/validation/models/h3_provider.py b/rl_engine/validation/models/h3_provider.py new file mode 100644 index 000000000..3eec1818f --- /dev/null +++ b/rl_engine/validation/models/h3_provider.py @@ -0,0 +1,192 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The diffusers MiniMax-H3 conditioning path, op for op, without diffusers. + +Each function replays the tensor ops of the pinned reference +(``huggingface/diffusers@f53d552``, see ``h3_manifest.json``) in the same +order and dtypes, so it runs the same kernels the provider does on a given +device. This is the *provider* side of the RFC #420 comparisons; the FP32/FP64 +goldens live in ``rl_engine.reference.minimax_h3``. +""" + +from __future__ import annotations + +import math + +import torch +import torch.nn.functional as F + + +def provider_time_proj(timestep: torch.Tensor, num_channels: int = 256) -> torch.Tensor: + """Replay provider FP32 cosine/sine features for ``(T,)`` timesteps. + + Matches ``Timesteps(256, flip_sin_to_cos=True, downscale_freq_shift=0)`` + with the default channel count, returning shape ``(T, 256)``. + """ + + half_dim = num_channels // 2 + exponent = -math.log(10000) * torch.arange( + start=0, end=half_dim, dtype=torch.float32, device=timestep.device + ) + exponent = exponent / (half_dim - 0) + emb = torch.exp(exponent) + emb = timestep[:, None].float() * emb[None, :] + emb = 1 * emb + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) + return emb + + +def provider_time_embedder( + features: torch.Tensor, + w1: torch.Tensor, + b1: torch.Tensor, + w2: torch.Tensor, + b2: torch.Tensor, +) -> torch.Tensor: + """Replay linear_1, SiLU, and linear_2 using the checkpoint weight dtype. + + H3's FP32 weights map features of shape ``(T, 256)`` to ``(T, 2688)``. + """ + + sample = F.linear(features.to(w1.dtype), w1, b1) + sample = F.silu(sample) + return F.linear(sample, w2, b2) + + +def provider_adaln_modulation( + temb: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor, hidden_size: int = 5376 +) -> tuple[torch.Tensor, ...]: + """Replay provider SiLU, the weight-dtype cast, and the AdaLN projection. + + Return six ``(3T, hidden_size)`` views in timestep-major modality order + for checkpoint weights shaped ``(18 * hidden_size, 2688)``. + """ + + temb = F.linear(F.silu(temb).to(weight.dtype), weight, bias) + temb = temb.view(-1, 6 * hidden_size) + return temb.chunk(6, dim=-1) + + +def provider_adaln_row_gather( + modulation: tuple[torch.Tensor, ...], timestep_indices: torch.Tensor, token_tags: torch.Tensor +) -> tuple[torch.Tensor, ...]: + """``adaln_indices = timestep_indices * 3 + token_tags`` then six ``index_select``.""" + + adaln_indices = timestep_indices * 3 + token_tags + return tuple(tensor.index_select(0, adaln_indices) for tensor in modulation) + + +def provider_norm_modulate( + hidden_states: torch.Tensor, + norm_weight: torch.Tensor, + shift: torch.Tensor, + scale: torch.Tensor, + indices: torch.Tensor, + eps: float = 1e-5, +) -> torch.Tensor: + """``norm(x) * (1.0 + scale[i]) + shift[i]`` with ``index_select``; block and norm_out.""" + + norm_hidden_states = F.rms_norm(hidden_states, (hidden_states.shape[-1],), norm_weight, eps) + return norm_hidden_states * (1.0 + scale.index_select(0, indices)) + shift.index_select( + 0, indices + ) + + +def provider_gate_residual( + residual: torch.Tensor, gate: torch.Tensor, indices: torch.Tensor, sublayer_output: torch.Tensor +) -> torch.Tensor: + """``residual + gate.index_select(0, adaln_indices) * sublayer_output`` (block).""" + + return residual + gate.index_select(0, indices) * sublayer_output + + +def provider_final_adaln_out( + hidden_states: torch.Tensor, + norm_weight: torch.Tensor, + temb: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor, + timestep_indices: torch.Tensor, + eps: float = 1e-5, +) -> torch.Tensor: + """``MiniMaxH3AdaLayerNormOut.forward``: shift first, indexed by timestep.""" + + shift, scale = F.linear(F.silu(temb).to(weight.dtype), weight, bias).chunk(2, dim=-1) + hidden_states = F.rms_norm(hidden_states, (hidden_states.shape[-1],), norm_weight, eps) + return hidden_states * (1.0 + scale.index_select(0, timestep_indices)) + shift.index_select( + 0, timestep_indices + ) + + +# --- One transformer block (RFC #420 ``ws1_one_h3_block``) ------------------------------------ + + +def provider_rope( + position_ids: torch.Tensor, rope_freq_dim: int = 16, rope_theta: float = 10000.0 +) -> tuple[torch.Tensor, torch.Tensor]: + """``MiniMaxH3RotaryPosEmbed``: FP32 ``(S, 96)`` cos/sin over the ``(t, h, w)`` axes.""" + + inv_freq = 1.0 / ( + rope_theta + ** ( + torch.arange(0, 2 * rope_freq_dim, 2, dtype=torch.float32, device=position_ids.device) + / (2 * rope_freq_dim) + ) + ) + position_ids = position_ids.to(torch.float32) + freqs = position_ids.unsqueeze(-1) * inv_freq.view(1, 1, -1) + freqs_t, freqs_h, freqs_w = freqs.unbind(dim=1) + freqs = torch.cat((freqs_t, freqs_h, freqs_w), dim=-1) + freqs = torch.cat((freqs, freqs), dim=-1) + return freqs.cos(), freqs.sin() + + +def provider_apply_rotary( + hidden_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor +) -> torch.Tensor: + """``_apply_rotary_emb``: rotate-half on the leading 96 of 128 channels of ``(B, S, H, D)``. + + The FP32 tables are cast to the activation dtype (BF16) before the rotation. + """ + + rotary_dim = cos.shape[-1] + hidden_states_rotary = hidden_states[..., :rotary_dim] + hidden_states_pass = hidden_states[..., rotary_dim:] + cos = cos.to(hidden_states.dtype)[None, :, None, :] + sin = sin.to(hidden_states.dtype)[None, :, None, :] + x1, x2 = hidden_states_rotary.chunk(2, dim=-1) + hidden_states_rotated = torch.cat((-x2, x1), dim=-1) + hidden_states_rotary = hidden_states_rotary * cos + hidden_states_rotated * sin + return torch.cat((hidden_states_rotary, hidden_states_pass), dim=-1).contiguous() + + +def provider_linear(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + """A bias-free ``nn.Linear`` of the block (Q/K/V, O and both FFN projections).""" + + return F.linear(x, weight) + + +def provider_head_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float = 1e-5) -> torch.Tensor: + """``norm_q``/``norm_k``: ``nn.RMSNorm(128)`` on ``(B, S, H, D)``.""" + + return F.rms_norm(x, (x.shape[-1],), weight, eps) + + +def provider_attention(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor) -> torch.Tensor: + """The native diffusers backend: SDPA on ``(B, H, S, D)`` views, non-causal, no mask. + + Returns ``(B, S, H*D)`` after ``flatten(2, 3).type_as(query)``. + """ + + q, k, v = (x.permute(0, 2, 1, 3) for x in (query, key, value)) + out = F.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False) + return out.permute(0, 2, 1, 3).flatten(2, 3).type_as(query) + + +def provider_swiglu(projected: torch.Tensor) -> torch.Tensor: + """diffusers ``SwiGLU`` after its projection: ``hidden * silu(gate)`` on the two halves.""" + + hidden_states, gate = projected.chunk(2, dim=-1) + return hidden_states * F.silu(gate) diff --git a/rl_engine/validation/models/h3_report.py b/rl_engine/validation/models/h3_report.py new file mode 100644 index 000000000..7ac7e9f08 --- /dev/null +++ b/rl_engine/validation/models/h3_report.py @@ -0,0 +1,806 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Per-operator performance and accuracy measurements for the H3 evidence. + +``PERF_CASES[op]`` builds the timed cases that ``benchmarks/models/benchmark_h3_conditioning.py`` +prints and ``tools/validation/models/h3_evidence.py`` stores. ``ACCURACY[op]`` measures the +op against the provider path and its golden. Each RFC #420 row registers both. +""" + +from __future__ import annotations + +import statistics +from typing import Any, Callable + +import torch + +from rl_engine.runtime.registry import KernelRegistry +from rl_engine.validation.models.h3_cases import h3_packed_layout, h3_timesteps +from rl_engine.validation.models.h3_provider import ( + provider_adaln_modulation, + provider_adaln_row_gather, + provider_final_adaln_out, + provider_gate_residual, + provider_norm_modulate, + provider_time_embedder, + provider_time_proj, +) +from rl_engine.validation.models.h3_weights import ( + WEIGHTS_ENV, + h3_weights_dir, + load_h3_conditioning_weights, +) + +# Keys of a perf case that hold a timed callable, in display order. +TIMED_KEYS = ( + "candidate", + "candidate_checked", + "provider", + "candidate_backward", + "provider_backward", +) + + +def _prepared_call( + fn: Callable[..., Any], setup: Callable[[], Any] | None = None +) -> Callable[[], Any]: + if setup is None: + return fn + prepared = setup() + return lambda: fn(prepared) + + +def _sample_us(fn: Callable[..., Any], setup: Callable[[], Any] | None = None) -> float: + run = _prepared_call(fn, setup) + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + run() + end.record() + end.synchronize() + return start.elapsed_time(end) * 1e3 + + +def time_us( + fn: Callable[..., Any], + warmup: int = 20, + iters: int = 200, + *, + setup: Callable[[], Any] | None = None, +) -> float: + """Median CUDA-event microseconds; optional setup runs outside the timed region. + + When supplied, ``setup`` runs before every call and its return value is passed + to ``fn``. Backward measurements use it to build a fresh forward graph. + """ + + for _ in range(warmup): + _prepared_call(fn, setup)() + torch.cuda.synchronize() + return statistics.median(_sample_us(fn, setup) for _ in range(iters)) + + +def peak_mib(fn: Callable[..., Any], *, setup: Callable[[], Any] | None = None) -> float: + """Incremental peak CUDA memory of the call, excluding optional setup.""" + + run = _prepared_call(fn, setup) + torch.cuda.synchronize() + torch.cuda.reset_peak_memory_stats() + base = torch.cuda.memory_allocated() + run() + torch.cuda.synchronize() + return (torch.cuda.max_memory_allocated() - base) / 2**20 + + +def measure(case: dict[str, Any], warmup: int = 20, iters: int = 200) -> dict[str, Any]: + """Interleave timed keys, reversing their execution order on alternate iterations. + + Record raw CUDA-event samples, executed orders, median microseconds, + decimal GB/s from ``case['bytes']``, and peak allocated MiB per callable. + """ + + row = {key: case[key] for key in ("op", "case", "backend", "bytes") if key in case} + keys = [key for key in TIMED_KEYS if key in case] + if not keys: + return row + orders = (keys, list(reversed(keys))) + row["execution_order"] = { + "policy": "alternating", + "iteration_0": orders[0], + "iteration_1": orders[1], + } + if any(key.endswith("_backward") for key in keys): + row["backward_timing_scope"] = "backward_only" + for iteration in range(warmup): + for key in orders[iteration % 2]: + _prepared_call(case[key], case.get(f"{key}_setup"))() + torch.cuda.synchronize() + samples = {key: [] for key in keys} + for iteration in range(iters): + for key in orders[iteration % 2]: + samples[key].append(_sample_us(case[key], case.get(f"{key}_setup"))) + row["timing_samples_us"] = samples + row["timing_order"] = [orders[iteration % 2] for iteration in range(iters)] + for key in keys: + us = statistics.median(samples[key]) + row[f"{key}_us"] = us + row[f"{key}_gbps"] = case["bytes"] / (us * 1e-6) / 1e9 + row[f"{key}_peak_mib"] = peak_mib(case[key], setup=case.get(f"{key}_setup")) + return row + + +# --------------------------------------------------------------------------- # +# timestep_sinusoid_h3 +# --------------------------------------------------------------------------- # + + +def _sinusoid_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + """Build CUDA sinusoid timing cases with checked, unchecked, and provider calls.""" + + op = registry.get_op("timestep_sinusoid_h3", device="cuda") + cases = [] + for num in (1, 2, 4, 64): + t = h3_timesteps(num) + cases.append( + { + "op": "timestep_sinusoid_h3", + "case": f"T={num}", + "backend": type(op).__name__, + "bytes": num * 4 + num * 256 * 4, + "candidate": lambda t=t: op.forward(t, check_range=False), + "candidate_checked": lambda t=t: op.forward(t), + "provider": lambda t=t: provider_time_proj(t), + } + ) + return cases + + +def _sinusoid_accuracy(registry: KernelRegistry) -> dict[str, Any]: + """Compare sinusoid outputs to provider and FP64 golden across timestep counts.""" + + op = registry.get_op("timestep_sinusoid_h3", device="cuda") + golden = registry._get_or_create_backend( + registry._priority_map["cpu"]["timestep_sinusoid_h3"][-1] + ) + rows = [] + for num in (1, 2, 3, 7, 64, 1000, 4097): + t = h3_timesteps(num, seed=num) + out, ref, gold = op(t), provider_time_proj(t), golden.forward_fp32(t) + rows.append( + { + "num_timesteps": num, + "bitwise_equal_to_provider": bool(torch.equal(out, ref)), + "max_abs_vs_fp64": float((out - gold).abs().max()), + "provider_max_abs_vs_fp64": float((ref - gold).abs().max()), + } + ) + return {"contract_atol": 1e-5, "cases": rows} + + +def h3_params(names: list[str], shapes: list[tuple[int, ...]]) -> list[torch.Tensor]: + """Return CUDA checkpoint parameters, or seeded random tensors when weights are unset.""" + + if h3_weights_dir() is not None: + weights = load_h3_conditioning_weights("cuda", names) + return [weights[name] for name in names] + generator = torch.Generator(device="cuda").manual_seed(0) + return [ + torch.randn(shape, device="cuda", generator=generator) * shape[-1] ** -0.5 + for shape in shapes + ] + + +# --------------------------------------------------------------------------- # +# timestep_mlp_fp32 +# --------------------------------------------------------------------------- # + +MLP_NAMES = [f"time_embedder.linear_{i}.{p}" for i in (1, 2) for p in ("weight", "bias")] +MLP_SHAPES = [(5376, 256), (5376,), (2688, 5376), (2688,)] + + +def _mlp_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + """Build CUDA MLP candidate/provider timing cases for one to four timesteps.""" + + op = registry.get_op("timestep_mlp_fp32", device="cuda") + params = h3_params(MLP_NAMES, MLP_SHAPES) + weight_bytes = sum(p.numel() * p.element_size() for p in params) + cases = [] + for num in (1, 2, 3, 4): + x = torch.rand(num, 256, device="cuda") * 2 - 1 + cases.append( + { + "op": "timestep_mlp_fp32", + "case": f"T={num}", + "backend": type(op).__name__, + "bytes": weight_bytes, + "candidate": lambda x=x: op(x, *params), + "provider": lambda x=x: provider_time_embedder(x, *params), + } + ) + return cases + + +def _mlp_accuracy(registry: KernelRegistry, draws: int = 200) -> dict[str, Any]: + """Measure per-draw MLP error against FP64 and test single-row/batched equality.""" + + op = registry.get_op("timestep_mlp_fp32", device="cuda") + golden = registry._get_or_create_backend(registry._priority_map["cpu"]["timestep_mlp_fp32"][-1]) + sinusoid = registry.get_op("timestep_sinusoid_h3", device="cuda") + params = h3_params(MLP_NAMES, MLP_SHAPES) + cuda_err, provider_err = [], [] + for seed in range(draws): + x = sinusoid(h3_timesteps(4, seed=seed)) + gold = golden.forward_fp32(x, *params) + cuda_err.append(float((op(x, *params) - gold).abs().max())) + provider_err.append(float((provider_time_embedder(x, *params) - gold).abs().max())) + # Row invariance: each timestep alone equals its row in the batch of 9. + x = sinusoid(h3_timesteps(9, seed=1)) + full = op(x, *params) + invariant = all(torch.equal(op(x[i : i + 1], *params)[0], full[i]) for i in range(9)) + return { + "contract_atol": 1e-4, + "draws": draws, + "num_timesteps": 4, + "cuda_max_abs_vs_fp64": cuda_err, + "provider_max_abs_vs_fp64": provider_err, + "rows_batch_invariant": invariant, + } + + +# --------------------------------------------------------------------------- # +# adaln_projection_3mod +# --------------------------------------------------------------------------- # + +ADALN_NAMES = [ + "transformer_blocks.0.adaln_proj.linear.weight", + "transformer_blocks.0.adaln_proj.linear.bias", +] + + +def _adaln_params() -> tuple[torch.Tensor, torch.Tensor]: + """Return block-0 AdaLN projection weights and bias as BF16 CUDA tensors.""" + + weight, bias = h3_params(ADALN_NAMES, [(96768, 2688), (96768,)]) + return weight.bfloat16(), bias.bfloat16() + + +def _projection_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + """Build CUDA AdaLN candidate/provider timing cases for one to four timesteps.""" + + op = registry.get_op("adaln_projection_3mod", device="cuda") + weight, bias = _adaln_params() + cases = [] + for num in (1, 2, 3, 4): + temb = torch.randn(num, 2688, device="cuda") + cases.append( + { + "op": "adaln_projection_3mod", + "case": f"T={num}", + "backend": type(op).__name__, + "bytes": weight.numel() * weight.element_size(), + "candidate": lambda temb=temb: op(temb, weight, bias), + "provider": lambda temb=temb: provider_adaln_modulation(temb, weight, bias), + } + ) + return cases + + +def _flat_cat(outputs) -> torch.Tensor: + """Concatenate flattened modulation tensors in their output order.""" + + return torch.cat([out.reshape(-1) for out in outputs]) + + +def _projection_accuracy(registry: KernelRegistry, draws: int = 20) -> dict[str, Any]: + """Report per-draw BF16 equality to FP64 goldens and single-row/batch invariance.""" + + op = registry.get_op("adaln_projection_3mod", device="cuda") + golden = registry._get_or_create_backend( + registry._priority_map["cpu"]["adaln_projection_3mod"][-1] + ) + weight, bias = _adaln_params() + cuda_frac, provider_frac, early_frac = [], [], [] + for seed in range(draws): + g = torch.Generator(device="cuda").manual_seed(seed) + temb = torch.randn(4, 2688, device="cuda", generator=g) * 2 + gold = _flat_cat(golden.forward_fp32(temb, weight, bias)).bfloat16() + early = _flat_cat(golden.forward_fp32(temb.bfloat16().float(), weight, bias)).bfloat16() + ours = _flat_cat(op(temb, weight, bias)) + theirs = _flat_cat(provider_adaln_modulation(temb, weight, bias)) + cuda_frac.append(float((ours == gold).float().mean())) + provider_frac.append(float((theirs == gold).float().mean())) + early_frac.append(float((ours == early).float().mean())) + temb = torch.randn(9, 2688, device="cuda") + full = op(temb, weight, bias) + invariant = all( + all( + torch.equal(s, f[3 * i : 3 * i + 3]) + for s, f in zip(op(temb[i : i + 1], weight, bias), full, strict=True) + ) + for i in range(9) + ) + return { + "draws": draws, + "num_timesteps": 4, + "cuda_correctly_rounded": cuda_frac, + "provider_correctly_rounded": provider_frac, + "early_cast_golden_match": early_frac, + "rows_batch_invariant": invariant, + } + + +# --------------------------------------------------------------------------- # +# adaln_row_gather +# --------------------------------------------------------------------------- # + +GATHER_SEQ_LENS = (4097, 32768, 131072) + + +def _gather_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + op = registry.get_op("adaln_row_gather", device="cuda") + num_timesteps, hidden = 3, 5376 + rows = torch.randn(3 * num_timesteps, 6 * hidden, device="cuda").bfloat16() + chunks = rows.chunk(6, dim=-1) + cases = [] + for seq in GATHER_SEQ_LENS: + ti, tags = h3_packed_layout(seq, num_timesteps, seed=seq) + grads = [torch.randn(seq, hidden, device="cuda").bfloat16() for _ in range(6)] + + def prepare_backward(fn, ti=ti, tags=tags): + leaf = rows.detach().requires_grad_(True) + return list(fn(leaf, ti, tags)) + + def backward(outputs, grads=grads): + torch.autograd.backward(outputs, grads) + + cases.append( + { + "op": "adaln_row_gather", + "case": f"S={seq}", + "backend": type(op).__name__, + # bytes written (six (S, H) BF16 outputs); the 3T-row table stays in L2 + "bytes": 6 * seq * hidden * 2, + "candidate": lambda ti=ti, tags=tags: op.forward(rows, ti, tags, check_range=False), + "provider": lambda ti=ti, tags=tags: provider_adaln_row_gather(chunks, ti, tags), + "candidate_backward": backward, + "provider_backward": backward, + "candidate_backward_setup": lambda b=prepare_backward: b( + lambda r, indices, row_tags: op.forward(r, indices, row_tags, check_range=False) + ), + "provider_backward_setup": lambda b=prepare_backward: b( + lambda r, indices, row_tags: provider_adaln_row_gather( + r.chunk(6, dim=-1), indices, row_tags + ) + ), + } + ) + return cases + + +def _gather_accuracy(registry: KernelRegistry) -> dict[str, Any]: + from rl_engine.validation.models.h3_chain import run_backward_case + + op = registry.get_op("adaln_row_gather", device="cuda") + golden = registry._get_or_create_backend(registry._priority_map["cpu"]["adaln_row_gather"][-1]) + rows = torch.randn(9, 6 * 5376, device="cuda").bfloat16() + forward_bitwise = {} + for seq in (1, 257, 4097, 32768): + ti, tags = h3_packed_layout(seq, 3, seed=seq) + ours = op(rows, ti, tags) + theirs = provider_adaln_row_gather(rows.chunk(6, dim=-1), ti, tags) + forward_bitwise[str(seq)] = all( + torch.equal(a, b) for a, b in zip(ours, theirs, strict=True) + ) + + # Op-level backward: correctly rounded FP32 segment sums vs the atomic BF16 scatter-add. + ti, tags = h3_packed_layout(4097, 3, seed=1) + grads = [torch.randn(4097, 5376, device="cuda").bfloat16() for _ in range(6)] + + def grad_of(fn): + leaf = rows.detach().clone().requires_grad_(True) + torch.autograd.backward(list(fn(leaf)), grads) + return leaf.grad + + ref = rows.detach().double().requires_grad_(True) + torch.autograd.backward(list(golden.forward_fp32(ref, ti, tags)), [g.float() for g in grads]) + gold = ref.grad.bfloat16() + backward = {} + for name, fn in ( + ("cuda", lambda leaf: op(leaf, ti, tags)), + ("provider", lambda leaf: provider_adaln_row_gather(leaf.chunk(6, dim=-1), ti, tags)), + ): + first, second = grad_of(fn), grad_of(fn) + backward[name] = { + "repeat_bitwise_equal": bool(torch.equal(first, second)), + "correctly_rounded_fraction": float((first == gold).float().mean()), + } + + result: dict[str, Any] = { + "forward_bitwise_vs_index_select": forward_bitwise, + "op_backward": backward, + "chain_backward": [], + } + if h3_weights_dir() is None: + result["chain_backward_skipped"] = ( + f"{WEIGHTS_ENV} not set; run tools/weights/prepare_h3_weights.py for chain measurements" + ) + else: + weights = load_h3_conditioning_weights("cuda") + result["chain_backward"] = [ + run_backward_case(registry, weights, num_timesteps=t, seq_len=s) + for t, s in ((1, 257), (3, 257), (1, 4097), (3, 4097), (4, 32768)) + ] + return result + + +# --------------------------------------------------------------------------- # +# h3_rmsnorm (block norm1 + MSA modulation) +# --------------------------------------------------------------------------- # + +NORM_SEQ_LENS = (4097, 32768, 131072) + + +def _norm_inputs(seq: int, seed: int = 0): + weight = h3_params(["transformer_blocks.0.norm1.weight"], [(5376,)])[0].bfloat16() + g = torch.Generator(device="cuda").manual_seed(seed) + table = (torch.randn(9, 6 * 5376, device="cuda", generator=g) * 0.5).bfloat16() + shift, scale = table.view(9, 6, 5376)[:, 0], table.view(9, 6, 5376)[:, 1] + ti, tags = h3_packed_layout(seq, 3, seed=seed) + x = (torch.randn(1, seq, 5376, device="cuda", generator=g) * 2).bfloat16() + return x, weight, shift, scale, ti * 3 + tags + + +def _norm_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + op = registry.get_op("h3_rmsnorm", device="cuda") + cases = [] + for seq in NORM_SEQ_LENS: + x, weight, shift, scale, index = _norm_inputs(seq) + grad = torch.randn_like(x) + + def prepare_backward(fn, x=x, weight=weight, shift=shift, scale=scale): + leaves = [t.detach().requires_grad_(True) for t in (x, weight, shift, scale)] + return fn(*leaves) + + def backward(output, grad=grad): + output.backward(grad) + + cases.append( + { + "op": "h3_rmsnorm", + "case": f"S={seq}", + "backend": type(op).__name__, + "bytes": 2 * seq * 5376 * 2, # read x, write y + "candidate": lambda x=x, w=weight, sh=shift, sc=scale, i=index: ( + op.forward_modulated(x, w, sh, sc, i, check_range=False) + ), + "provider": lambda x=x, w=weight, sh=shift, sc=scale, i=index: ( + provider_norm_modulate(x, w, sh, sc, i) + ), + "candidate_backward": backward, + "provider_backward": backward, + "candidate_backward_setup": lambda b=prepare_backward, i=index: b( + lambda x_, w_, sh_, sc_: op.forward_modulated( + x_, w_, sh_, sc_, i, check_range=False + ) + ), + "provider_backward_setup": lambda b=prepare_backward, i=index: b( + lambda x_, w_, sh_, sc_: provider_norm_modulate(x_, w_, sh_, sc_, i) + ), + } + ) + return cases + + +def _norm_accuracy(registry: KernelRegistry) -> dict[str, Any]: + op = registry.get_op("h3_rmsnorm", device="cuda") + weight_source = "pinned_checkpoint" if h3_weights_dir() is not None else "synthetic" + names = [ + "transformer_blocks.0.norm1.weight", + "transformer_blocks.0.norm2.weight", + "token_refiner.final_norm.weight", + "norm_out.norm.weight", + ] + weights = { + name: weight.bfloat16() + for name, weight in zip(names, h3_params(names, [(5376,)] * len(names)), strict=True) + } + plain = {} + for name in names: + x = torch.randn(2, 777, 5376, device="cuda").bfloat16() + ref = torch.nn.functional.rms_norm(x, (5376,), weights[name], 1e-5) + plain[name] = bool(torch.equal(op(x, weights[name]), ref)) + x, weight, shift, scale, index = _norm_inputs(4097, seed=1) + modulated = bool( + torch.equal( + op.forward_modulated(x, weight, shift, scale, index), + provider_norm_modulate(x, weight, shift, scale, index), + ) + ) + invariant = bool( + torch.equal( + op.forward_modulated(x[:, 100:140], weight, shift, scale, index[100:140])[0], + op.forward_modulated(x, weight, shift, scale, index)[0, 100:140], + ) + ) + + # Backward against FP64, for the CUDA op and the diffusers expression. + grad = torch.randn( + x.shape, device="cuda", generator=torch.Generator(device="cuda").manual_seed(7) + ).to(x.dtype) + + def grads(fn, dtype=None): + tensors = [t if dtype is None else t.to(dtype) for t in (x, weight, shift, scale)] + leaves = [t.detach().clone().requires_grad_(True) for t in tensors] + fn(*leaves).backward(grad if dtype is None else grad.to(dtype)) + return [leaf.grad for leaf in leaves] + + def golden(x_, w_, sh_, sc_): + n = x_ * torch.rsqrt(x_.square().mean(-1, keepdim=True) + 1e-5) * w_ + return n * (1 + sc_.index_select(0, index)) + sh_.index_select(0, index) + + ref = grads(golden, torch.float64) + backward = {} + for name, fn in ( + ("cuda", lambda *t: op.forward_modulated(*t, index)), + ("provider", lambda *t: provider_norm_modulate(*t, index)), + ): + first, second = grads(fn), grads(fn) + backward[name] = { + "repeat_bitwise_equal": all( + torch.equal(a, b) for a, b in zip(first, second, strict=True) + ), + "rel_error": { + key: float((g.double() - r).abs().max() / r.abs().max()) + for key, g, r in zip(("dx", "dweight", "dshift", "dscale"), first, ref, strict=True) + }, + } + return { + "weight_source": weight_source, + "plain_bitwise_vs_nn_rmsnorm": plain, + "modulated_bitwise_vs_diffusers": modulated, + "rows_batch_invariant": invariant, + "backward": backward, + } + + +# --------------------------------------------------------------------------- # +# adaln_gate_residual (gate_msa after attention) +# --------------------------------------------------------------------------- # + + +def _gate_inputs(seq: int, seed: int = 0, dtype=torch.bfloat16): + g = torch.Generator(device="cuda").manual_seed(seed) + table = (torch.randn(9, 6 * 5376, device="cuda", generator=g) * 0.5).to(dtype) + residual = torch.randn(1, seq, 5376, device="cuda", generator=g).to(dtype) + y = (torch.randn(1, seq, 5376, device="cuda", generator=g) * 3).to(dtype) + ti, tags = h3_packed_layout(seq, 3, seed=seed) + return table, residual, y, ti * 3 + tags + + +def _gate_view(table): + return table.view(table.shape[0], 6, -1)[:, 2] + + +def _gate_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + op = registry.get_op("adaln_gate_residual", device="cuda") + cases = [] + for seq in NORM_SEQ_LENS: + table, residual, y, index = _gate_inputs(seq) + gate = _gate_view(table) + grad = torch.randn_like(residual) + + def prepare_backward(fn, residual=residual, y=y, table=table, grad=grad): + leaves = [t.detach().requires_grad_(True) for t in (residual, y, table)] + return fn(leaves[0], leaves[1], _gate_view(leaves[2])), grad + + def backward(state): + output, grad = state + output.backward(grad) + + cases.append( + { + "op": "adaln_gate_residual", + "case": f"S={seq}", + "backend": type(op).__name__, + "bytes": 3 * seq * 5376 * 2, # read residual and y, write out + "candidate": lambda r=residual, y=y, g=gate, i=index: op.forward( + r, y, g, i, check_range=False + ), + "provider": lambda r=residual, y=y, g=gate, i=index: provider_gate_residual( + r, g, i, y + ), + "candidate_backward": backward, + "provider_backward": backward, + "candidate_backward_setup": lambda b=prepare_backward, i=index: b( + lambda r, y_, g: op.forward(r, y_, g, i, check_range=False) + ), + "provider_backward_setup": lambda b=prepare_backward, i=index: b( + lambda r, y_, g: provider_gate_residual(r, g, i, y_) + ), + } + ) + return cases + + +def _gate_accuracy(registry: KernelRegistry) -> dict[str, Any]: + op = registry.get_op("adaln_gate_residual", device="cuda") + forward_bitwise = {} + for dtype in (torch.bfloat16, torch.float16, torch.float32): + table, residual, y, index = _gate_inputs(777, seed=1, dtype=dtype) + gate = _gate_view(table) + forward_bitwise[str(dtype).removeprefix("torch.")] = bool( + torch.equal( + op(residual, y, gate, index), provider_gate_residual(residual, gate, index, y) + ) + ) + table, residual, y, index = _gate_inputs(4097, seed=2) + invariant = bool( + torch.equal( + op(residual[:, 10:60], y[:, 10:60], _gate_view(table), index[10:60])[0], + op(residual, y, _gate_view(table), index)[0, 10:60], + ) + ) + grad = torch.randn( + residual.shape, device="cuda", generator=torch.Generator(device="cuda").manual_seed(7) + ).to(residual.dtype) + + def grads(fn, dtype=None): + tensors = [t if dtype is None else t.to(dtype) for t in (residual, y, table)] + leaves = [t.detach().clone().requires_grad_(True) for t in tensors] + fn(leaves[0], leaves[1], _gate_view(leaves[2])).backward( + grad if dtype is None else grad.to(dtype) + ) + return [leaf.grad for leaf in leaves] + + ref = grads(lambda r, y_, g: provider_gate_residual(r, g, index, y_), torch.float64) + backward = {} + for name, fn in ( + ("cuda", lambda r, y_, g: op(r, y_, g, index)), + ("provider", lambda r, y_, g: provider_gate_residual(r, g, index, y_)), + ): + first, second = grads(fn), grads(fn) + backward[name] = { + "repeat_bitwise_equal": all( + torch.equal(a, b) for a, b in zip(first, second, strict=True) + ), + "rel_error": { + key: float((g.double() - r).abs().max() / r.abs().max()) + for key, g, r in zip( + ("d_residual", "d_sublayer", "d_gate"), first, ref, strict=True + ) + }, + } + return { + "forward_bitwise_vs_diffusers": forward_bitwise, + "rows_batch_invariant": invariant, + "backward": backward, + } + + +# --------------------------------------------------------------------------- # +# final_adaln_out (norm_out) +# --------------------------------------------------------------------------- # + +FINAL_NAMES = ["norm_out.norm.weight", "norm_out.linear.weight", "norm_out.linear.bias"] + + +def _final_inputs(seq: int, seed: int = 0): + norm_weight, weight, bias = h3_params(FINAL_NAMES, [(5376,), (10752, 2688), (10752,)]) + g = torch.Generator(device="cuda").manual_seed(seed) + x = (torch.randn(1, seq, 5376, device="cuda", generator=g) * 2).bfloat16() + temb = torch.randn(3, 2688, device="cuda", generator=g) * 2 + ti, _ = h3_packed_layout(seq, 3, seed=seed) + return x, norm_weight.bfloat16(), temb, weight.bfloat16(), bias.bfloat16(), ti + + +def _final_perf(registry: KernelRegistry) -> list[dict[str, Any]]: + op = registry.get_op("final_adaln_out", device="cuda") + cases = [] + for seq in NORM_SEQ_LENS: + x, nw, temb, w, b, ti = _final_inputs(seq) + grad = torch.randn_like(x) + + def prepare_backward(fn, tensors=(x, nw, temb, w, b), grad=grad): + leaves = [t.detach().requires_grad_(True) for t in tensors] + return fn(*leaves), grad + + def backward(state): + output, grad = state + output.backward(grad) + + cases.append( + { + "op": "final_adaln_out", + "case": f"S={seq}", + "backend": type(op).__name__, + # read x + norm_out.linear, write out + "bytes": 2 * seq * 5376 * 2 + w.numel() * 2, + "candidate": lambda t=(x, nw, temb, w, b), i=ti: op(*t, i), + "provider": lambda t=(x, nw, temb, w, b), i=ti: provider_final_adaln_out(*t, i), + "candidate_backward": backward, + "provider_backward": backward, + "candidate_backward_setup": lambda bw=prepare_backward, i=ti: bw( + lambda *t: op(*t, i) + ), + "provider_backward_setup": lambda bw=prepare_backward, i=ti: bw( + lambda *t: provider_final_adaln_out(*t, i) + ), + } + ) + return cases + + +def _final_accuracy(registry: KernelRegistry) -> dict[str, Any]: + op = registry.get_op("final_adaln_out", device="cuda") + golden_op = registry._get_or_create_backend( + registry._priority_map["cpu"]["final_adaln_out"][-1] + ) + x, nw, temb, w, b, ti = _final_inputs(4097, seed=1) + ours = op(x, nw, temb, w, b, ti) + theirs = provider_final_adaln_out(x, nw, temb, w, b, ti) + golden = golden_op.forward_fp32(x, nw, temb, w, b, ti) + invariant = bool( + torch.equal(op(x[:, 100:160], nw, temb, w, b, ti[100:160])[0], ours[0, 100:160]) + ) + grad = torch.randn( + x.shape, device="cuda", generator=torch.Generator(device="cuda").manual_seed(7) + ).to(x.dtype) + + def grads(fn, dtype=None): + tensors = [t if dtype is None else t.to(dtype) for t in (x, nw, temb, w, b)] + leaves = [t.detach().clone().requires_grad_(True) for t in tensors] + fn(*leaves).backward(grad if dtype is None else grad.to(dtype)) + return [leaf.grad for leaf in leaves] + + def golden_backward(x_, nw_, t_, w_, b_): + # Keep the declared BF16 boundaries with FP64 leaves and identity VJPs. + act = t_ * torch.sigmoid(t_) + act = act + (act.to(torch.bfloat16).double() - act).detach() + table = torch.nn.functional.linear(act, w_, b_) + table = table + (table.to(torch.bfloat16).double() - table).detach() + shift, scale = table.chunk(2, dim=-1) + n = x_ * torch.rsqrt(x_.square().mean(-1, keepdim=True) + 1e-5) * nw_ + return n * (1.0 + scale.index_select(0, ti)) + shift.index_select(0, ti) + + ref = grads(golden_backward, torch.float64) + backward = {} + for name, fn in ( + ("cuda", lambda *t: op(*t, ti)), + ("provider", lambda *t: provider_final_adaln_out(*t, ti)), + ): + first, second = grads(fn), grads(fn) + backward[name] = { + "repeat_bitwise_equal": all(torch.equal(a, c) for a, c in zip(first, second)), + "rel_error": { + key: float((g.double() - r).abs().max() / r.abs().max()) + for key, g, r in zip(("dx", "d_norm_w", "d_temb", "dW", "db"), first, ref) + }, + } + return { + "equal_to_diffusers_fraction": float((ours == theirs).float().mean()), + "max_abs_vs_golden": float((ours.float() - golden).abs().max()), + "provider_max_abs_vs_golden": float((theirs.float() - golden).abs().max()), + "rows_batch_invariant": invariant, + "backward": backward, + } + + +PERF_CASES: dict[str, Callable[[KernelRegistry], list[dict[str, Any]]]] = { + "timestep_sinusoid_h3": _sinusoid_perf, + "timestep_mlp_fp32": _mlp_perf, + "adaln_projection_3mod": _projection_perf, + "adaln_row_gather": _gather_perf, + "h3_rmsnorm": _norm_perf, + "adaln_gate_residual": _gate_perf, + "final_adaln_out": _final_perf, +} +ACCURACY: dict[str, Callable[[KernelRegistry], dict[str, Any]]] = { + "timestep_sinusoid_h3": _sinusoid_accuracy, + "timestep_mlp_fp32": _mlp_accuracy, + "adaln_projection_3mod": _projection_accuracy, + "adaln_row_gather": _gather_accuracy, + "h3_rmsnorm": _norm_accuracy, + "adaln_gate_residual": _gate_accuracy, + "final_adaln_out": _final_accuracy, +} diff --git a/rl_engine/validation/models/h3_weights.py b/rl_engine/validation/models/h3_weights.py new file mode 100644 index 000000000..f587f2726 --- /dev/null +++ b/rl_engine/validation/models/h3_weights.py @@ -0,0 +1,108 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Pinned MiniMax-H3 conditioning weights for the RFC #420 tests. + +Only the tensors listed in ``h3_manifest.json`` are used. They all live in +shard 1 of 14, so ``tools/weights/prepare_h3_weights.py`` downloads that shard, +checks it against the manifest and extracts the tensors into +``$RL_KERNEL_H3_WEIGHTS/h3_conditioning.safetensors``. With ``--block`` it also +extracts block 0's attention and FFN weights (manifest ``block_tensors``, +about 1.1 GB) into ``h3_block0.safetensors`` for ``ws1_one_h3_block``. +""" + +from __future__ import annotations + +import hashlib +import json +import os +from pathlib import Path +from typing import Any + +import torch + +MANIFEST_PATH = Path(__file__).with_name("h3_manifest.json") +WEIGHTS_ENV = "RL_KERNEL_H3_WEIGHTS" +EXTRACTED_FILE = "h3_conditioning.safetensors" +BLOCK_FILE = "h3_block0.safetensors" + + +def load_h3_manifest() -> dict[str, Any]: + """Read the bundled checkpoint identity and pinned tensor contracts.""" + + return json.loads(MANIFEST_PATH.read_text()) + + +def h3_weights_dir() -> Path | None: + """Expand the configured extraction directory, or return None when unset or blank.""" + + value = os.environ.get(WEIGHTS_ENV, "").strip() + return Path(value).expanduser() if value else None + + +def sha256_file(path: Path, chunk: int = 1 << 24) -> str: + """Return a file SHA256 while reading bounded chunks.""" + + digest = hashlib.sha256() + with path.open("rb") as handle: + while block := handle.read(chunk): + digest.update(block) + return digest.hexdigest() + + +def sha256_tensor(tensor: torch.Tensor) -> str: + """Hash a CPU tensor's contiguous bytes without safetensors metadata.""" + + tensor_bytes = tensor.contiguous().view(torch.uint8).numpy() + return hashlib.sha256(memoryview(tensor_bytes)).hexdigest() + + +def _load_pinned( + filename: str, section: str, device: torch.device | str, names: list[str] | None +) -> dict[str, torch.Tensor]: + """Load one extracted file and check every tensor of a manifest section against it.""" + + from safetensors.torch import load_file + + root = h3_weights_dir() + if root is None: + raise FileNotFoundError(f"set {WEIGHTS_ENV} to the prepare_h3_weights.py output dir") + path = root / filename + if not path.is_file(): + flag = " --block" if section == "block_tensors" else "" + raise FileNotFoundError(f"{path} missing; run tools/weights/prepare_h3_weights.py{flag}") + specs = load_h3_manifest()[section] + tensors = load_file(str(path), device="cpu") + wanted = names if names is not None else list(specs) + for name, spec in specs.items(): + tensor = tensors[name] + if str(tensor.dtype).removeprefix("torch.") != spec["dtype"]: + raise ValueError(f"{name}: dtype {tensor.dtype} != manifest {spec['dtype']}") + if list(tensor.shape) != spec["shape"]: + raise ValueError(f"{name}: shape {list(tensor.shape)} != manifest {spec['shape']}") + actual = sha256_tensor(tensor) + if actual != spec["sha256"]: + raise ValueError(f"{name}: sha256 {actual} != manifest {spec['sha256']}") + out: dict[str, torch.Tensor] = {} + for name in wanted: + if name not in specs: + raise KeyError(name) + out[name] = tensors[name].to(device=device) + return out + + +def load_h3_conditioning_weights( + device: torch.device | str = "cpu", names: list[str] | None = None +) -> dict[str, torch.Tensor]: + """Load the extracted pinned tensors; raise if they are missing or off-manifest.""" + + return _load_pinned(EXTRACTED_FILE, "tensors", device, names) + + +def load_h3_block_weights(device: torch.device | str = "cpu") -> dict[str, torch.Tensor]: + """Conditioning tensors plus block 0's attention/FFN weights, every one sha-checked.""" + + return { + **load_h3_conditioning_weights(device), + **_load_pinned(BLOCK_FILE, "block_tensors", device, None), + } diff --git a/rl_engine/validation/models/h3_ws2.py b/rl_engine/validation/models/h3_ws2.py new file mode 100644 index 000000000..7342774a0 --- /dev/null +++ b/rl_engine/validation/models/h3_ws2.py @@ -0,0 +1,418 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""WS1 vs WS2 comparisons for the H3 conditioning path (tests and evidence). + +The per-rank helpers run inside real multi-process groups (tests and +``tools/validation/models/h3_ws2_evidence.py``); the WS1 side runs on one GPU. +""" + +from __future__ import annotations + +from typing import Any + +import torch + + +def ws1_projection(temb, weight, bias, grad) -> dict[str, torch.Tensor]: + """The single-GPU ``adaln_projection_3mod`` table and its three gradients.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection import ( + H3AdaLNProjectionCudaOp, + ) + + leaves = [t.detach().clone().requires_grad_(True) for t in (temb, weight, bias)] + table = H3AdaLNProjectionCudaOp().forward_table(*leaves) + table.backward(grad) + return { + "table": table.detach(), + "d_temb": leaves[0].grad, + "d_weight": leaves[1].grad, + "d_bias": leaves[2].grad, + } + + +def tp_projection_rank(collective, temb, weight, bias, grad) -> dict[str, Any]: + """One TP rank: shard the full weight, run the TP op, return its view of the results.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.tp_adaln_projection import ( + H3TPAdaLNProjectionCudaOp, + shard_adaln_projection, + ) + + op = H3TPAdaLNProjectionCudaOp(collective, weight.shape[0]) + w_shard, b_shard = shard_adaln_projection(weight, bias, collective.world_size, collective.rank) + leaves = [t.detach().clone().requires_grad_(True) for t in (temb, w_shard, b_shard)] + table = op.forward_table(*leaves) + table.backward(grad) + return { + "table": table.detach(), + "d_temb": leaves[0].grad, + "d_weight_shard": leaves[1].grad, + "d_bias_shard": leaves[2].grad, + "readback": op.readback(), + } + + +def tp_matches_ws1(ws1: dict[str, torch.Tensor], ranks: list[dict[str, Any]]) -> dict[str, bool]: + """Byte equality of every rank's replicated outputs and of the concatenated shards.""" + + return { + "table": all(torch.equal(r["table"], ws1["table"]) for r in ranks), + "d_temb": all(torch.equal(r["d_temb"], ws1["d_temb"]) for r in ranks), + "d_weight": torch.equal(torch.cat([r["d_weight_shard"] for r in ranks]), ws1["d_weight"]), + "d_bias": torch.equal(torch.cat([r["d_bias_shard"] for r in ranks]), ws1["d_bias"]), + } + + +def tp_projection_sweep(collective, temb, weight, bias, grad, iters: int = 100) -> dict[str, Any]: + """Evidence for one rank: per-T byte equality against WS1 on this GPU, timing, readback. + + ``temb``/``grad`` hold the largest T; every T in ``1..T`` uses their first + rows. Each rank compares its own replicated outputs and its own shard of + ``dW``/``db`` with the WS1 result computed here, so nothing large is shipped. + """ + + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection import ( + H3AdaLNProjectionCudaOp, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.tp_adaln_projection import ( + H3TPAdaLNProjectionCudaOp, + adaln_column_shard, + ) + from rl_engine.validation.models.h3_report import time_us + + shard = adaln_column_shard(weight.shape[0], collective.world_size, collective.rank) + cols = slice(shard.begin, shard.end) + tp_op, ws1_op = ( + H3TPAdaLNProjectionCudaOp(collective, weight.shape[0]), + H3AdaLNProjectionCudaOp(), + ) + equality, perf = [], [] + for num_t in range(1, temb.shape[0] + 1): + t, g = temb[:num_t], grad[:num_t] + ws1 = ws1_projection(t, weight, bias, g) + mine = tp_projection_rank(collective, t, weight, bias, g) + equality.append( + { + "num_timesteps": num_t, + "table": torch.equal(mine["table"], ws1["table"]), + "d_temb": torch.equal(mine["d_temb"], ws1["d_temb"]), + "d_weight_shard": torch.equal(mine["d_weight_shard"], ws1["d_weight"][cols]), + "d_bias_shard": torch.equal(mine["d_bias_shard"], ws1["d_bias"][cols]), + } + ) + leaves = [x.detach().clone().requires_grad_(True) for x in (t, weight[cols], bias[cols])] + full = [x.detach().clone().requires_grad_(True) for x in (t, weight, bias)] + + def step(op, args): + for x in args: + x.grad = None + op.forward_table(*args).backward(g) + + perf.append( + { + "num_timesteps": num_t, + "tp_forward_us": time_us( + lambda: tp_op.forward_table(t, weight[cols], bias[cols]), 30, iters + ), + "tp_forward_backward_us": time_us(lambda: step(tp_op, leaves), 10, iters), + "ws1_forward_us": time_us(lambda: ws1_op.forward_table(t, weight, bias), 30, iters), + "ws1_forward_backward_us": time_us(lambda: step(ws1_op, full), 10, iters), + } + ) + return {"equality": equality, "perf": perf, "readback": tp_op.readback()} + + +def _region(gate_fn, norm_fn, residual, y, weight, table): + """One block's MLP-side SP region: ``norm2(residual + gate_msa * y)`` with modulation.""" + + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = table.chunk(6, dim=-1) + hidden = gate_fn(residual, y, gate_msa) + return norm_fn(hidden, weight, shift_mlp, scale_mlp) + + +def ws1_sp_region(residual, y, weight, table, index, grad) -> dict[str, torch.Tensor]: + """The SP region on one GPU with the WS1 ops (the reference for ``sp_norm_adaln``).""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual import H3GateResidualCudaOp + from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import H3RMSNormCudaOp + + gate_op, norm_op = H3GateResidualCudaOp(), H3RMSNormCudaOp() + leaves = [t.detach().clone().requires_grad_(True) for t in (residual, y, weight, table)] + out = _region( + lambda r, yy, g: gate_op(r, yy, g, index), + lambda h, w, sh, sc: norm_op.forward_modulated(h, w, sh, sc, index), + *leaves, + ) + out.backward(grad) + names = ("d_residual", "d_y", "d_weight", "d_table") + return {"out": out.detach(), **{n: leaf.grad for n, leaf in zip(names, leaves)}} + + +def sp_region_rank(collective, residual, y, weight, table, index, grad) -> dict[str, Any]: + """One SP rank: its rows of the region and its (replicated) parameter gradients.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import H3SPNormAdaLNCudaOp + + op = H3SPNormAdaLNCudaOp(collective, seq_len=residual.shape[1], batch=residual.shape[0]) + rows = slice(op.layout.lo, op.layout.hi) + local = [t[:, rows].contiguous() for t in (residual, y)] + leaves = [t.detach().clone().requires_grad_(True) for t in (*local, weight, table)] + out = _region( + lambda r, yy, g: op.gate_residual(r, yy, g, index), + lambda h, w, sh, sc: op.norm_modulated(h, w, sh, sc, index), + *leaves, + ) + out.backward(grad[:, rows].contiguous()) + names = ("d_residual", "d_y", "d_weight", "d_table") + return { + "out": out.detach(), + **{n: leaf.grad for n, leaf in zip(names, leaves)}, + "readback": op.readback(), + } + + +def sp_matches_ws1(ws1: dict[str, torch.Tensor], ranks: list[dict[str, Any]]) -> dict[str, bool]: + """Rows: each rank's slice equals WS1's; parameters: every rank's copy equals WS1's.""" + + def rows(key): + return all( + torch.equal(r[key], ws1[key][:, slice(*r["readback"]["positions"])]) for r in ranks + ) + + return { + "out": rows("out"), + "d_residual": rows("d_residual"), + "d_y": rows("d_y"), + "d_weight": all(torch.equal(r["d_weight"], ws1["d_weight"]) for r in ranks), + "d_table": all(torch.equal(r["d_table"], ws1["d_table"]) for r in ranks), + } + + +def sp_case(batch, seq, hidden=5376, num_t=3, *, layout="block", dtype=torch.bfloat16, seed=0): + """Inputs of the SP region; identical on every rank (CPU generator).""" + + from rl_engine.validation.models.h3_cases import h3_block_layout, h3_packed_layout + + g = torch.Generator(device="cpu").manual_seed(seed) + make = h3_block_layout if layout == "block" else h3_packed_layout + ti, tags = make(seq, num_t, seed=seed) + return { + "residual": torch.randn(batch, seq, hidden, generator=g).to(dtype).cuda(), + "y": (torch.randn(batch, seq, hidden, generator=g) * 3).to(dtype).cuda(), + "weight": (torch.rand(hidden, generator=g) + 0.5).to(dtype).cuda(), + "table": (torch.randn(3 * num_t, 6 * hidden, generator=g) * 0.5).to(dtype).cuda(), + "index": (ti * 3 + tags).cuda(), + "grad": torch.randn(batch, seq, hidden, generator=g).to(dtype).cuda(), + } + + +def naive_sp_mismatch(case: dict[str, torch.Tensor], sp: int) -> dict[str, float]: + """Fraction of parameter-gradient elements where a naive SP backward differs from WS1. + + Naive: each rank runs the WS1 backward on its own rows, then the per-rank results + are summed in rank order (what a plain all-reduce of local gradients does). + """ + + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import sp_row_layout + + ws1 = ws1_sp_region(**case) + parts = [] + for rank in range(sp): + lay = sp_row_layout(case["residual"].shape[1], case["residual"].shape[0], sp, rank) + part = slice(lay.lo, lay.hi) + local = {k: (v[:, part] if k in ("residual", "y", "grad") else v) for k, v in case.items()} + local["index"] = case["index"][part] + parts.append(ws1_sp_region(**local)) + out = {} + for key in ("d_weight", "d_table"): + total = parts[0][key].float() + for p in parts[1:]: + total = total + p[key].float() + out[key] = (total.to(ws1[key].dtype) != ws1[key]).float().mean().item() + return out + + +def sp_region_sweep(collective, cases: list[dict[str, Any]], iters: int = 20) -> dict[str, Any]: + """Evidence for one SP rank: byte equality, rows exchanged and time, per case.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import _segment_tiles + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + _dweight_tiles, + _Plan, + sp_row_layout, + ) + + results = [] + for spec in cases: + case = sp_case(**spec) + batch, seq = case["residual"].shape[:2] + ws1 = ws1_sp_region(**case) + mine = sp_region_rank(collective, **case) + equal = sp_matches_ws1(ws1, [mine]) + lay = sp_row_layout(seq, batch, collective.world_size, collective.rank) + tiles = _segment_tiles(case["index"].repeat(batch), case["table"].shape[0]) + plan = _Plan(lay, [_dweight_tiles(batch * seq, "cuda"), tuple(tiles[:3])], "cuda") + del ws1, mine + torch.cuda.empty_cache() + results.append( + { + **spec, + "equal": equal, + "local_rows": batch * lay.local_len, + "rows_sent": int(plan.send_local.numel()), + "rows_gathered_per_rank": plan.max_send, + **_region_times(collective, case, lay, iters), + } + ) + return {"cases": results} + + +def _region_times(collective, case, lay, iters) -> dict[str, float]: + """Forward + backward time of the region (leaves prepared once, grads reset per call).""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual import H3GateResidualCudaOp + from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import H3RMSNormCudaOp + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import H3SPNormAdaLNCudaOp + from rl_engine.validation.models.h3_report import time_us + + index, rows = case["index"], slice(lay.lo, lay.hi) + sp_op = H3SPNormAdaLNCudaOp(collective, seq_len=lay.seq_len, batch=lay.batch) + gate_op, norm_op = H3GateResidualCudaOp(), H3RMSNormCudaOp() + + def leaves(local): + acts = [case[k][:, rows] if local else case[k] for k in ("residual", "y")] + return [ + t.detach().contiguous().requires_grad_(True) + for t in (*acts, case["weight"], case["table"]) + ] + + def step(fns, args, grad): + for leaf in args: + leaf.grad = None + _region(*fns, *args).backward(grad) + + sp_fns = ( + lambda r, yy, g: sp_op.gate_residual(r, yy, g, index), + lambda h, w, sh, sc: sp_op.norm_modulated(h, w, sh, sc, index), + ) + ws1_fns = ( + lambda r, yy, g: gate_op(r, yy, g, index), + lambda h, w, sh, sc: norm_op.forward_modulated(h, w, sh, sc, index), + ) + sp_args, ws1_args = leaves(True), leaves(False) + sp_grad = case["grad"][:, rows].contiguous() + return { + "sp_forward_backward_us": time_us(lambda: step(sp_fns, sp_args, sp_grad), 3, iters), + "ws1_forward_backward_us": time_us(lambda: step(ws1_fns, ws1_args, case["grad"]), 3, iters), + } + + +# --- real multi-process groups ------------------------------------------------- + + +def _bootstrap(rank, world, init_method, queue, inputs_path, target, kwargs): + import traceback + + import torch.distributed as dist + + from rl_engine.distributed.algorithms.collectives import DeterministicCollective + + try: + torch.cuda.set_device(rank) + dist.init_process_group( + "nccl", init_method=init_method, rank=rank, world_size=world, device_id=rank + ) + inputs = {k: v.cuda() for k, v in torch.load(inputs_path).items()} + collective = DeterministicCollective(device=rank, max_size_bytes=256 << 20) + out = target(collective, **inputs, **kwargs) + # Results go through a file: tensors sent over the queue are shared by + # file descriptor and vanish if this process exits before the parent reads. + torch.save( + {k: v.cpu() if isinstance(v, torch.Tensor) else v for k, v in out.items()}, + inputs_path.with_name(f"rank{rank}.pt"), + ) + collective.close() + dist.destroy_process_group() + queue.put({"rank": rank}) + except Exception: # forwarded to the parent + queue.put({"rank": rank, "error": traceback.format_exc()}) + + +def run_world(world: int, target, inputs: dict[str, torch.Tensor], timeout: float = 900, **kwargs): + """Run ``target(collective, **inputs, **kwargs)`` on ranks ``0..world-1`` (one GPU each). + + ``target`` must be importable (it is pickled into spawned processes). Each + rank gets a NCCL group and a ``DeterministicCollective``; results come back + on the CPU, ordered by rank. Any rank's exception is re-raised here. + """ + + import tempfile + import time + from pathlib import Path + from queue import Empty + + import torch.multiprocessing as mp + + if torch.cuda.device_count() < world: + raise RuntimeError(f"needs {world} GPUs, found {torch.cuda.device_count()}") + ctx = mp.get_context("spawn") + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "inputs.pt" + torch.save({k: v.cpu() for k, v in inputs.items()}, path) + queue = ctx.Queue() + init = (Path(tmp) / "init").as_uri() + procs = [ + ctx.Process(target=_bootstrap, args=(r, world, init, queue, path, target, kwargs)) + for r in range(world) + ] + started = [] + deadline = time.monotonic() + timeout + status = {} + try: + for proc in procs: + proc.start() + started.append(proc) + while len(status) < world or any(proc.is_alive() for proc in procs): + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"rank results timed out after {timeout}s") + try: + result = queue.get(timeout=min(0.1, remaining)) + except Empty: + exited = [ + (rank, proc.exitcode) + for rank, proc in enumerate(procs) + if proc.exitcode is not None and (proc.exitcode != 0 or rank not in status) + ] + if not exited: + continue + # A worker may have flushed its final message between the + # timed read and the exit-code check. + try: + result = queue.get_nowait() + except Empty: + rank, code = exited[0] + raise RuntimeError( + f"rank {rank} exited with code {code} without completing" + ) from None + if "error" in result: + raise RuntimeError(f"rank {result['rank']} failed:\n{result['error']}") + status[result["rank"]] = result + for rank, proc in enumerate(procs): + proc.join() + if proc.exitcode != 0: + raise RuntimeError(f"rank {rank} exited with code {proc.exitcode}") + return [{"rank": r, **torch.load(path.with_name(f"rank{r}.pt"))} for r in range(world)] + finally: + for proc in started: + if proc.is_alive(): + proc.terminate() + for proc in started: + proc.join(timeout=5) + if proc.is_alive(): + proc.kill() + proc.join() + queue.close() + queue.join_thread() diff --git a/rl_engine/validation/operators/operator_inputs.py b/rl_engine/validation/operators/operator_inputs.py index ca3b7c120..48856ebe0 100644 --- a/rl_engine/validation/operators/operator_inputs.py +++ b/rl_engine/validation/operators/operator_inputs.py @@ -16,6 +16,10 @@ DEFAULT_VOCAB = 151936 DEFAULT_ROPE_THETA = 1.0e6 DEFAULT_RMS_EPS = 1.0e-6 +H3_FREQ_DIM = 256 +H3_TIME_HIDDEN = 5376 +H3_TIME_EMBED = 2688 +H3_HIDDEN = 5376 def make_operator_inputs( @@ -24,6 +28,7 @@ def make_operator_inputs( dtype: torch.dtype, device: torch.device, ) -> dict[str, Any]: + """Build operator keyword inputs on the requested device from CLI shape and seed options.""" builders = { "rms_norm": _make_rms_norm_inputs, "qk_norm": _make_qk_norm_inputs, @@ -42,6 +47,13 @@ def make_operator_inputs( "embedding": _make_embedding_inputs, "lm_head": _make_lm_head_inputs, "kv_cache_attention": _make_kv_cache_attention_inputs, + "timestep_sinusoid_h3": _make_timestep_sinusoid_h3_inputs, + "timestep_mlp_fp32": _make_timestep_mlp_fp32_inputs, + "adaln_projection_3mod": _make_adaln_projection_3mod_inputs, + "adaln_row_gather": _make_adaln_row_gather_inputs, + "h3_rmsnorm": _make_h3_rmsnorm_inputs, + "adaln_gate_residual": _make_adaln_gate_residual_inputs, + "final_adaln_out": _make_final_adaln_out_inputs, } try: return builders[op_name](args, dtype, device) @@ -50,6 +62,7 @@ def make_operator_inputs( def operator_shape_name(op_name: str, args: argparse.Namespace) -> str: + """Return the operator's dimension label for benchmark and evidence reports.""" batch, seq = _batch_seq(args) vocab = _arg_int(args, "vocab", DEFAULT_VOCAB) names = { @@ -72,6 +85,16 @@ def operator_shape_name(op_name: str, args: argparse.Namespace) -> str: "embedding": f"{batch}x{seq}x{vocab}x{_normalized_dim(args)}", "lm_head": f"{batch}x{seq}x{_normalized_dim(args)}x{vocab}", "kv_cache_attention": f"{batch}x{DEFAULT_N_HEADS}x1x{seq + 1}x{DEFAULT_HEAD_DIM}", + "timestep_sinusoid_h3": f"{_h3_num_timesteps(args)}x{H3_FREQ_DIM}", + "timestep_mlp_fp32": f"{_h3_num_timesteps(args)}x{H3_FREQ_DIM}x{H3_TIME_HIDDEN}" + f"x{H3_TIME_EMBED}", + "adaln_projection_3mod": f"{_h3_num_timesteps(args)}x{H3_TIME_EMBED}" + f"x{6 * 3 * _h3_hidden(args)}", + "adaln_row_gather": f"{3 * _h3_num_timesteps(args)}x{6 * _h3_hidden(args)}" + f"->{batch * seq}", + "h3_rmsnorm": f"{batch}x{seq}x{_h3_hidden(args)}", + "adaln_gate_residual": f"{batch}x{seq}x{_h3_hidden(args)}", + "final_adaln_out": f"{batch}x{seq}x{_h3_hidden(args)}", } try: return names[op_name] @@ -331,6 +354,151 @@ def _make_kv_cache_attention_inputs( } +def _h3_num_timesteps(args: argparse.Namespace) -> int: + """Read the packed H3 timestep count, falling back to the CLI batch size.""" + # H3 packs a handful of distinct timesteps; reuse --batch as their count. + return _arg_int(args, "num_timesteps", _arg_int(args, "batch", 2)) + + +def _h3_timesteps( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> torch.Tensor: + """Build seeded timesteps in [0, 1] with endpoint cases in the requested storage dtype.""" + num = _h3_num_timesteps(args) + mode = _arg_str(args, "input_mode", "random") + if mode == "constant": + t = torch.full((num,), 0.5, device=device) + else: + t = torch.rand((num,), generator=_generator(args, device, offset=7), device=device) + t[0] = 0.0 + if num > 1: + t[-1] = 1.0 + return t.to(dtype) + + +def _make_timestep_sinusoid_h3_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + """Supply a one-dimensional packed timestep tensor to the H3 sinusoid operator.""" + return {"timestep": _h3_timesteps(args, dtype, device)} + + +def _make_timestep_mlp_fp32_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + """Build [T, 256] sinusoid features and seeded FP32 256→5376→2688 MLP parameters.""" + # The H3 time_embedder is declared FP32; ``dtype`` does not apply to it. + del dtype + from rl_engine.reference.minimax_h3.timestep_sinusoid import NativeH3TimestepSinusoidOp + + features = NativeH3TimestepSinusoidOp().forward(_h3_timesteps(args, torch.float32, device)) + scale1, scale2 = H3_FREQ_DIM**-0.5, H3_TIME_HIDDEN**-0.5 + return { + "x": features, + "w1": _floating_tensor((H3_TIME_HIDDEN, H3_FREQ_DIM), args, torch.float32, device, 1) + * scale1, + "b1": _floating_tensor((H3_TIME_HIDDEN,), args, torch.float32, device, 2) * 0.1, + "w2": _floating_tensor((H3_TIME_EMBED, H3_TIME_HIDDEN), args, torch.float32, device, 3) + * scale2, + "b2": _floating_tensor((H3_TIME_EMBED,), args, torch.float32, device, 4) * 0.1, + } + + +def _h3_hidden(args: argparse.Namespace) -> int: + """Read the AdaLN channel width, defaulting to the checkpoint's 5376 channels.""" + return _arg_int(args, "normalized_dim", H3_HIDDEN) + + +def _make_adaln_projection_3mod_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + """Build FP32 [T, 2688] embeddings and 18H projection parameters in the chosen dtype.""" + # temb is FP32 by contract (the SiLU runs before the cast); ``dtype`` is the + # projection's weight dtype, BF16 in the checkpoint. + n_out = 6 * 3 * _h3_hidden(args) + temb = _floating_tensor( + (_h3_num_timesteps(args), H3_TIME_EMBED), args, torch.float32, device, 0 + ) + weight = _floating_tensor((n_out, H3_TIME_EMBED), args, torch.float32, device, 1) + bias = _floating_tensor((n_out,), args, torch.float32, device, 2) + return { + "temb": temb * 2.0, + "weight": (weight * H3_TIME_EMBED**-0.5).to(dtype), + "bias": (bias * 0.1).to(dtype), + } + + +def _make_adaln_row_gather_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + # T distinct timesteps (--batch) and a packed sequence of batch * seq rows + # mixing all three modalities. + batch, seq = _batch_seq(args) + num_timesteps = _h3_num_timesteps(args) + generator = _generator(args, device, offset=29) + packed = batch * seq + token_tags = torch.randint(0, 3, (packed,), generator=generator, device=device) + timestep_indices = torch.randint( + 0, num_timesteps, (packed,), generator=generator, device=device + ) + rows = _floating_tensor((3 * num_timesteps, 6 * _h3_hidden(args)), args, dtype, device, 0) + return {"rows": rows, "timestep_indices": timestep_indices, "token_tags": token_tags} + + +def _make_h3_rmsnorm_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + batch, seq = _batch_seq(args) + hidden = _h3_hidden(args) + weight = _floating_tensor((hidden,), args, torch.float32, device, 1).abs() + 0.5 + return { + "x": _floating_tensor((batch, seq, hidden), args, dtype, device, 0), + "weight": weight.to(dtype), + "eps": 1e-5, + } + + +def _make_adaln_gate_residual_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + # gate rows: 3 modality rows for each of --num-timesteps (default --batch) timesteps. + batch, seq = _batch_seq(args) + hidden = _h3_hidden(args) + rows = 3 * _h3_num_timesteps(args) + generator = _generator(args, device, offset=31) + return { + "residual": _floating_tensor((batch, seq, hidden), args, dtype, device, 0), + "y": _floating_tensor((batch, seq, hidden), args, dtype, device, 1), + "gate": (_floating_tensor((rows, hidden), args, torch.float32, device, 2) * 0.5).to(dtype), + "index": torch.randint(0, rows, (seq,), generator=generator, device=device), + } + + +def _make_final_adaln_out_inputs( + args: argparse.Namespace, dtype: torch.dtype, device: torch.device +) -> dict[str, Any]: + # temb is FP32 by contract; norm_out.linear/norm/x use ``dtype`` (BF16 in H3). + batch, seq = _batch_seq(args) + hidden = _h3_hidden(args) + num_timesteps = _h3_num_timesteps(args) + # Realistic scales (norm weight in [0.5, 1], unit activations): with unscaled + # inputs the parameter gradients sum hundreds of rows into values ~1e3, and + # their cancelling elements cannot meet a per-element FP32 atol. + linear = _floating_tensor((2 * hidden, H3_TIME_EMBED), args, torch.float32, device, 3) + norm = 0.5 + 0.5 * torch.sigmoid(_floating_tensor((hidden,), args, torch.float32, device, 4)) + generator = _generator(args, device, offset=37) + return { + "x": _floating_tensor((batch, seq, hidden), args, dtype, device, 0), + "norm_weight": norm.to(dtype), + "temb": _floating_tensor((num_timesteps, H3_TIME_EMBED), args, torch.float32, device, 1), + "weight": (linear * H3_TIME_EMBED**-0.5).to(dtype), + "bias": (_floating_tensor((2 * hidden,), args, torch.float32, device, 2) * 0.1).to(dtype), + "timestep_indices": torch.randint( + 0, num_timesteps, (seq,), generator=generator, device=device + ), + } + + def _floating_tensor( shape: tuple[int, ...], args: argparse.Namespace, diff --git a/rl_engine/validation/operators/operator_specs.py b/rl_engine/validation/operators/operator_specs.py index 2323f4787..619174762 100644 --- a/rl_engine/validation/operators/operator_specs.py +++ b/rl_engine/validation/operators/operator_specs.py @@ -262,6 +262,110 @@ def _load_object(path: str) -> Any: }, grad_input_names=("logits",), ), + # MiniMax-H3 (RFC #420) conditioning path. The op is FP32 end to end, so + # its FP32 rows are the declared contract; other dtypes only change the + # timestep input's storage dtype. + "timestep_sinusoid_h3": OperatorSpec( + name="timestep_sinusoid_h3", + op_class="elementwise", + gold_path="rl_engine.reference.minimax_h3.timestep_sinusoid.NativeH3TimestepSinusoidOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ( + "rl_engine.reference.minimax_h3.timestep_sinusoid.NativeH3TimestepSinusoidOp" + ), + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "timestep_sinusoid.H3TimestepSinusoidCudaOp" + ), + }, + grad_input_names=("timestep",), + ), + "timestep_mlp_fp32": OperatorSpec( + name="timestep_mlp_fp32", + op_class="reduction", + gold_path="rl_engine.reference.minimax_h3.timestep_mlp.NativeH3TimestepMLPOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": "rl_engine.reference.minimax_h3.timestep_mlp.NativeH3TimestepMLPOp", + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "timestep_mlp.H3TimestepMLPCudaOp" + ), + }, + grad_input_names=("x", "w1", "b1", "w2", "b2"), + ), + "adaln_projection_3mod": OperatorSpec( + name="adaln_projection_3mod", + op_class="reduction", + gold_path=("rl_engine.reference.minimax_h3.adaln_projection.NativeH3AdaLNProjectionOp"), + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ( + "rl_engine.reference.minimax_h3.adaln_projection.NativeH3AdaLNProjectionOp" + ), + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "adaln_projection.H3AdaLNProjectionCudaOp" + ), + }, + grad_input_names=("temb", "weight", "bias"), + ), + # The gather's forward is a copy (asserted bitwise in tests/models/minimax_h3), but its VJP + # sums every packed position of a table row, so it is judged as a reduction. + "adaln_row_gather": OperatorSpec( + name="adaln_row_gather", + op_class="reduction", + gold_path="rl_engine.reference.minimax_h3.adaln_row_gather.NativeH3AdaLNRowGatherOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.minimax_h3.adaln_row_gather.NativeH3AdaLNRowGatherOp"), + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "adaln_row_gather.H3AdaLNRowGatherCudaOp" + ), + }, + grad_input_names=("rows",), + ), + "h3_rmsnorm": OperatorSpec( + name="h3_rmsnorm", + op_class="reduction", + gold_path="rl_engine.reference.minimax_h3.rmsnorm.NativeH3RMSNormOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": "rl_engine.reference.minimax_h3.rmsnorm.NativeH3RMSNormOp", + "cuda": "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm.H3RMSNormCudaOp", + }, + grad_input_names=("x", "weight"), + ), + "adaln_gate_residual": OperatorSpec( + name="adaln_gate_residual", + op_class="reduction", # forward is elementwise; the gate VJP sums positions + gold_path="rl_engine.reference.minimax_h3.gate_residual.NativeH3GateResidualOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": "rl_engine.reference.minimax_h3.gate_residual.NativeH3GateResidualOp", + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "gate_residual.H3GateResidualCudaOp" + ), + }, + grad_input_names=("residual", "y", "gate"), + ), + "final_adaln_out": OperatorSpec( + name="final_adaln_out", + op_class="reduction", + gold_path="rl_engine.reference.minimax_h3.final_adaln_out.NativeH3FinalAdaLNOutOp", + gold_method="forward_fp32", + candidate_paths={ + "pytorch": ("rl_engine.reference.minimax_h3.final_adaln_out.NativeH3FinalAdaLNOutOp"), + "cuda": ( + "rl_engine.backends.cuda.model_specific.minimax_h3." + "final_adaln_out.H3FinalAdaLNOutCudaOp" + ), + }, + grad_input_names=("x", "norm_weight", "temb", "weight", "bias"), + ), "pack": OperatorSpec( name="pack", op_class="elementwise", diff --git a/tests/models/minimax_h3/conftest.py b/tests/models/minimax_h3/conftest.py new file mode 100644 index 000000000..c20e425de --- /dev/null +++ b/tests/models/minimax_h3/conftest.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Shared fixtures for the MiniMax-H3 (RFC #420) operator tests. + +Real-weight tests read the pinned tensors from ``$RL_KERNEL_H3_WEIGHTS`` +(written by ``tools/weights/prepare_h3_weights.py``) and skip when it is unset. +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.validation.models.h3_weights import ( + WEIGHTS_ENV, + h3_weights_dir, + load_h3_block_weights, + load_h3_conditioning_weights, +) + + +@pytest.fixture(scope="session") +def h3_weights_cpu() -> dict[str, torch.Tensor]: + """Load pinned CPU weights once per session, skipping unavailable artifacts.""" + + if h3_weights_dir() is None: + pytest.skip(f"{WEIGHTS_ENV} not set; run tools/weights/prepare_h3_weights.py") + try: + return load_h3_conditioning_weights("cpu") + except FileNotFoundError as exc: + pytest.skip(str(exc)) + + +@pytest.fixture(scope="session") +def h3_block_weights_cuda() -> dict[str, torch.Tensor]: + """Conditioning plus block-0 tensors on CUDA, skipping when they were not extracted.""" + + if h3_weights_dir() is None: + pytest.skip(f"{WEIGHTS_ENV} not set; run tools/weights/prepare_h3_weights.py --block") + if not torch.cuda.is_available(): + pytest.skip("needs CUDA") + try: + return load_h3_block_weights("cuda") + except FileNotFoundError as exc: + pytest.skip(str(exc)) diff --git a/tests/models/minimax_h3/test_h3_adaln_gate_residual.py b/tests/models/minimax_h3/test_h3_adaln_gate_residual.py new file mode 100644 index 000000000..e4c303e10 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_adaln_gate_residual.py @@ -0,0 +1,115 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``adaln_gate_residual``: ``residual + gate[row] * sublayer_output``. + +* forward and ``d_sublayer`` are bitwise equal to diffusers' eager + expression in every dtype (the gate row is gathered in the kernel); +* ``d_gate`` is a deterministic FP32 segment sum, checked against FP64; +* rows are elementwise-independent; malformed inputs fail closed. +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.reference.minimax_h3.gate_residual import NativeH3GateResidualOp +from rl_engine.validation.models.h3_cases import h3_packed_layout +from rl_engine.validation.models.h3_provider import provider_gate_residual + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +def _cuda_op(): + from rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual import ( + H3GateResidualCudaOp, + h3_gate_residual_available, + ) + + if not h3_gate_residual_available(): + pytest.skip("rl_engine._C lacks h3_gate_residual_*") + return H3GateResidualCudaOp() + + +def _case(seq, hidden=5376, dtype=torch.bfloat16, seed=0, batch=2): + g = torch.Generator(device="cpu").manual_seed(seed) + table = (torch.randn(9, 6 * hidden, generator=g) * 0.5).to(dtype).cuda() + residual = torch.randn(batch, seq, hidden, generator=g).to(dtype).cuda() + y = (torch.randn(batch, seq, hidden, generator=g) * 3).to(dtype).cuda() + ti, tags = h3_packed_layout(seq, 3, seed=seed) + return table, residual, y, ti * 3 + tags + + +def _gate(table, chunk=2): + return table.view(table.shape[0], 6, -1)[:, chunk] # gate_msa (2) / gate_mlp (5) view + + +class TestReference: + def test_rejects_bad_inputs(self): + op = NativeH3GateResidualOp() + res, y, gate = torch.zeros(4, 8), torch.zeros(4, 8), torch.zeros(3, 8) + with pytest.raises(IndexError): + op(res, y, gate, torch.tensor([0, 1, 2, 3])) + with pytest.raises(ValueError): + op(res, y[:, :4], gate, torch.tensor([0, 1, 2, 0])) + with pytest.raises(TypeError): + op(res, y, gate.bfloat16(), torch.tensor([0, 1, 2, 0])) + with pytest.raises(ValueError): # S mismatch + op(res, y, gate, torch.tensor([0, 1])) + + +@requires_cuda +class TestCuda: + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) + @pytest.mark.parametrize("seq, hidden", [(4097, 5376), (33, 12), (17, 13)]) + def test_forward_bitwise_equal_to_diffusers(self, dtype, seq, hidden): + table, residual, y, index = _case(seq, hidden, dtype, seed=seq) + for chunk in (2, 5): + ours = _cuda_op()(residual, y, _gate(table, chunk), index) + assert torch.equal( + ours, provider_gate_residual(residual, _gate(table, chunk), index, y) + ) + + def test_rows_are_batch_and_position_invariant(self): + table, residual, y, index = _case(500, seed=1) + full = _cuda_op()(residual, y, _gate(table), index) + part = _cuda_op()(residual[1:2, 100:150], y[1:2, 100:150], _gate(table), index[100:150]) + assert torch.equal(part[0], full[1, 100:150]) + + def test_backward(self): + table, residual, y, index = _case(4097, seed=2) + grad = torch.randn_like(residual) + + def grads(fn, dtype=None): + tensors = [t if dtype is None else t.to(dtype) for t in (residual, y, table)] + leaves = [t.detach().clone().requires_grad_(True) for t in tensors] + fn(leaves[0], leaves[1], _gate(leaves[2])).backward( + grad if dtype is None else grad.to(dtype) + ) + return [leaf.grad for leaf in leaves] + + ours = grads(lambda r, yy, gg: _cuda_op()(r, yy, gg, index)) + again = grads(lambda r, yy, gg: _cuda_op()(r, yy, gg, index)) + theirs = grads(lambda r, yy, gg: provider_gate_residual(r, gg, index, yy)) + golden = grads(lambda r, yy, gg: provider_gate_residual(r, gg, index, yy), torch.float64) + assert all(torch.equal(a, b) for a, b in zip(ours, again)) + assert torch.equal(ours[0], theirs[0]) # d_residual + assert torch.equal(ours[1], theirs[1]) # d_sublayer: exact product, one rounding + rel = ((ours[2].double() - golden[2]).abs().max() / golden[2].abs().max()).item() + assert rel < 5e-3 # only the BF16 rounding of the table gradient + assert torch.count_nonzero(ours[2].view(9, 6, -1)[:, [0, 1, 3, 4, 5]]) == 0 + + def test_rejects_cpu_and_bad_index(self): + table, residual, y, index = _case(10, hidden=16) + with pytest.raises(IndexError): + _cuda_op()(residual, y, _gate(table), index + 9) + with pytest.raises(ValueError): + _cuda_op()(residual.cpu(), y.cpu(), _gate(table).cpu(), index.cpu()) + + def test_registry_dispatches_cuda(self): + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("adaln_gate_residual", device="cuda") + assert type(op).__name__ == "H3GateResidualCudaOp" diff --git a/tests/models/minimax_h3/test_h3_adaln_modulation.py b/tests/models/minimax_h3/test_h3_adaln_modulation.py new file mode 100644 index 000000000..cd6647739 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_adaln_modulation.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Fused AdaLN modulation (projection + row gather) with an FP32 table gradient. + +The forward must be bitwise equal to the two RFC #420 ops run separately; the +backward must be deterministic and must not re-round the table gradient to +BF16 (so ``d_temb`` reaches FP64-golden accuracy). +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.validation.models.h3_cases import h3_packed_layout + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +def _ops(): + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_modulation import ( + H3AdaLNModulationCudaOp, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection import ( + H3AdaLNProjectionCudaOp, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + H3AdaLNRowGatherCudaOp, + adaln_row_gather_available, + ) + + if not adaln_row_gather_available(): + pytest.skip("rl_engine._C lacks the H3 kernels") + return H3AdaLNModulationCudaOp(), H3AdaLNProjectionCudaOp(), H3AdaLNRowGatherCudaOp() + + +def _params(hidden, dim, seed=0): + g = torch.Generator(device="cpu").manual_seed(seed) + n_out = 18 * hidden + weight = (torch.randn(n_out, dim, generator=g) / dim**0.5).bfloat16().cuda() + bias = (torch.randn(n_out, generator=g) * 0.1).bfloat16().cuda() + temb = (torch.randn(3, dim, generator=g) * 2).cuda() + return temb, weight, bias + + +def _grads(seq, hidden, seed=0): + g = torch.Generator(device="cuda").manual_seed(seed) + return [torch.randn(seq, hidden, device="cuda", generator=g).bfloat16() for _ in range(6)] + + +def _leaf_grads(fn, temb, weight, bias, grads): + leaves = [t.detach().clone().requires_grad_(True) for t in (temb, weight, bias)] + torch.autograd.backward(list(fn(*leaves)), grads) + return [leaf.grad for leaf in leaves] + + +@requires_cuda +class TestFusedModulation: + def test_forward_bitwise_equal_to_separate_ops(self): + fused, proj, gather = _ops() + temb, weight, bias = _params(64, 2688) + ti, tags = h3_packed_layout(1000, 3) + ours = fused(temb, weight, bias, ti, tags) + separate = gather.gather_chunks(proj(temb, weight, bias), ti, tags) + assert all(torch.equal(a, b) for a, b in zip(ours, separate)) + + def test_backward_keeps_table_gradient_fp32(self, h3_weights_cpu): + fused, proj, gather = _ops() + weight = h3_weights_cpu["transformer_blocks.0.adaln_proj.linear.weight"].cuda() + bias = h3_weights_cpu["transformer_blocks.0.adaln_proj.linear.bias"].cuda() + temb = torch.randn(3, 2688, device="cuda") * 2 + ti, tags = h3_packed_layout(4097, 3, seed=1) + grads = _grads(4097, 5376, seed=1) + ours = _leaf_grads(lambda t, w, b: fused(t, w, b, ti, tags), temb, weight, bias, grads) + separate = _leaf_grads( + lambda t, w, b: gather.gather_chunks(proj(t, w, b), ti, tags), + temb, + weight, + bias, + grads, + ) + + ref = [t.detach().double().requires_grad_(True) for t in (temb, weight, bias)] + act = ref[0] * torch.sigmoid(ref[0]) + act = act + (act.to(torch.bfloat16).double() - act).detach() # identity-VJP cast + table = F.linear(act, ref[1], ref[2]).view(-1, 6 * 5376) + index = ti * 3 + tags + torch.autograd.backward( + [chunk.index_select(0, index) for chunk in table.chunk(6, dim=-1)], + [g.double() for g in grads], + ) + + def rel(grad, golden): + return ((grad.double() - golden).abs().max() / golden.abs().max()).item() + + # d_temb: no BF16 rounding anywhere on its path, so it reaches FP32 accuracy. + assert rel(ours[0], ref[0].grad) < 1e-5 + assert rel(ours[0], ref[0].grad) < rel(separate[0], ref[0].grad) / 100 + # dW, db: only their own final BF16 rounding remains. + for grad, golden in zip(ours[1:], (ref[1].grad, ref[2].grad)): + assert (grad == golden.to(torch.bfloat16)).float().mean() > 0.9999 + + def test_backward_repeat_bitwise(self): + fused, _, _ = _ops() + temb, weight, bias = _params(64, 2688, seed=3) + ti, tags = h3_packed_layout(3000, 3, seed=3) + grads = _grads(3000, 64, seed=3) + first = _leaf_grads(lambda t, w, b: fused(t, w, b, ti, tags), temb, weight, bias, grads) + again = _leaf_grads(lambda t, w, b: fused(t, w, b, ti, tags), temb, weight, bias, grads) + assert all(torch.equal(a, b) for a, b in zip(first, again)) + + def test_rejects_bad_indices(self): + fused, _, _ = _ops() + temb, weight, bias = _params(8, 16) + ti, tags = h3_packed_layout(10, 3) + with pytest.raises(IndexError): + fused(temb, weight, bias, ti + 3, tags) + with pytest.raises(TypeError): + fused(temb.bfloat16(), weight, bias, ti, tags) diff --git a/tests/models/minimax_h3/test_h3_adaln_projection.py b/tests/models/minimax_h3/test_h3_adaln_projection.py new file mode 100644 index 000000000..6f9323ed5 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_adaln_projection.py @@ -0,0 +1,288 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``adaln_projection_3mod``: SiLU + 2688 -> 96768 projection, 6 x 3 table. + +* layout: the six outputs and the three modality rows per timestep land + exactly where diffusers puts them (RFC probe H4); +* mixed precision: SiLU in FP32, one cast to BF16; an early cast (probe H7) + is rejected at the API and is measurably different numerically; +* accuracy: forward and gradients against the declared-cast golden on the + pinned block-0 weights; +* invariance and determinism: a timestep's 18 modulation rows do not depend + on the other timesteps, and repeated runs are bitwise equal. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.reference.minimax_h3 import H3_ADALN_CHUNKS, H3_MODALITY_NUM +from rl_engine.reference.minimax_h3.adaln_projection import NativeH3AdaLNProjectionOp +from rl_engine.validation.models.h3_provider import provider_adaln_modulation + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +WEIGHT = "transformer_blocks.0.adaln_proj.linear.weight" +BIAS = "transformer_blocks.0.adaln_proj.linear.bias" + + +def _cuda_op(): + """Construct the CUDA projection operator, skipping missing native linear kernels.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection import ( + H3AdaLNProjectionCudaOp, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import det_linear_available + + if not det_linear_available(): + pytest.skip("rl_engine._C lacks h3_det_linear_*") + return H3AdaLNProjectionCudaOp() + + +def _temb(num, dim=2688, seed=0, device="cuda"): + """Generate seeded FP32 timestep embeddings on CPU and move them to the test device.""" + + g = torch.Generator(device="cpu").manual_seed(seed) + return (torch.randn(num, dim, generator=g) * 2).to(device) + + +def _synthetic(hidden, dim, dtype=torch.bfloat16, seed=0, device="cuda"): + """Create seeded projection weights for six chunks and three modalities per channel.""" + + g = torch.Generator(device="cpu").manual_seed(seed) + n_out = H3_ADALN_CHUNKS * H3_MODALITY_NUM * hidden + weight = (torch.randn(n_out, dim, generator=g) / dim**0.5).to(dtype) + bias = (torch.randn(n_out, generator=g) * 0.1).to(dtype) + return weight.to(device), bias.to(device) + + +def _flat(outputs): + """Flatten and concatenate all modulation chunks for elementwise comparisons.""" + + return torch.cat([out.reshape(-1) for out in outputs]) + + +class TestReference: + def test_provider_layout_matches_diffusers(self): + """Match the provider values and timestep-major modality layout on CPU.""" + + temb = _temb(2, dim=16, device="cpu") + weight, bias = _synthetic(8, 16, device="cpu") + ours = NativeH3AdaLNProjectionOp().forward(temb, weight, bias) + theirs = provider_adaln_modulation(temb, weight, bias, hidden_size=8) + assert len(ours) == 6 + for a, b in zip(ours, theirs, strict=True): + assert a.shape == (6, 8) and torch.equal(a, b) + + def test_rejects_bf16_temb(self): + # RFC probe H7: casting temb before the SiLU. + """Reject an early BF16 cast so SiLU receives the declared FP32 embeddings.""" + + weight, bias = _synthetic(8, 16, device="cpu") + with pytest.raises(TypeError, match="float32"): + NativeH3AdaLNProjectionOp().forward(_temb(2, 16, device="cpu").bfloat16(), weight, bias) + + def test_rejects_bad_shapes(self): + """Reject invalid projection dimensions, empty batches, and mismatched bias dtype.""" + + weight, bias = _synthetic(8, 16, device="cpu") + op = NativeH3AdaLNProjectionOp() + with pytest.raises(ValueError): # not a multiple of 6 * 3 rows + op.forward(_temb(2, 16, device="cpu"), weight[:-1], bias[:-1]) + with pytest.raises(ValueError): + op.forward(_temb(2, 15, device="cpu"), weight, bias) + with pytest.raises(ValueError): + op.forward(_temb(0, 16, device="cpu"), weight, bias) + with pytest.raises(TypeError): + op.forward(_temb(2, 16, device="cpu"), weight, bias.float()) + + +@requires_cuda +class TestLayout: + def test_channel_identity_lands_in_its_chunk_and_modality_row(self): + """Probe H4: bias o = m*6H + c*H + h tagged (m, c); zero weight passes it through.""" + + hidden, dim, num = 16, 32, 3 + n_out = H3_ADALN_CHUNKS * H3_MODALITY_NUM * hidden + o = torch.arange(n_out) + modality, chunk = o // (H3_ADALN_CHUNKS * hidden), (o // hidden) % H3_ADALN_CHUNKS + bias = (modality * 10 + chunk).to(torch.bfloat16).cuda() + weight = torch.zeros(n_out, dim, dtype=torch.bfloat16, device="cuda") + outputs = _cuda_op()(_temb(num, dim), weight, bias) + assert len(outputs) == 6 + for c, out in enumerate(outputs): + assert out.shape == (num * 3, hidden) + for t in range(num): + for m in range(3): + expected = torch.full((hidden,), float(m * 10 + c), device="cuda") + assert torch.equal(out[t * 3 + m].float(), expected), (c, t, m) + + def test_outputs_are_views_of_one_table_like_diffusers(self): + """Require all six modulation chunks to share storage with the expected offsets.""" + + weight, bias = _synthetic(16, 32) + outputs = _cuda_op()(_temb(2, 32), weight, bias) + base = outputs[0].untyped_storage().data_ptr() + assert all(out.untyped_storage().data_ptr() == base for out in outputs) + assert outputs[1].data_ptr() - outputs[0].data_ptr() == 16 * outputs[0].element_size() + + +@pytest.fixture(scope="module") +def block0(h3_weights_cpu): + """Move the pinned block-zero projection weight and bias to CUDA for this module.""" + + return h3_weights_cpu[WEIGHT].cuda(), h3_weights_cpu[BIAS].cuda() + + +@requires_cuda +class TestCudaRealWeights: + @pytest.mark.parametrize("num", [1, 2, 3, 4]) + def test_forward_is_correctly_rounded_declared_golden(self, block0, num): + """Check checkpoint accuracy and near-exact BF16 rounding of the declared golden.""" + + weight, bias = block0 + temb = _temb(num, seed=num) + outputs = _cuda_op()(temb, weight, bias) + golden = NativeH3AdaLNProjectionOp().forward_fp32(temb, weight, bias) + for out, gold in zip(outputs, golden, strict=True): + assert out.dtype == torch.bfloat16 and out.shape == (3 * num, 5376) + # tolerance_contract.json forward_accuracy / reduction / bfloat16. + torch.testing.assert_close(out.float(), gold, atol=5e-2, rtol=2e-2) + # Beyond the contract: almost every element is the correctly rounded + # golden value; the rest are 1-ULP ties of the FP32 accumulation. + match = (_flat(outputs) == _flat(golden).bfloat16()).float().mean().item() + assert match > 0.999 + + def test_early_cast_is_numerically_detectable(self, block0): + """Probe H7: a BF16 cast before the SiLU moves about half the outputs.""" + + weight, bias = block0 + temb = _temb(3, seed=7) + op = NativeH3AdaLNProjectionOp() + ours = _flat(_cuda_op()(temb, weight, bias)) + declared = _flat(op.forward_fp32(temb, weight, bias)).bfloat16() + early = _flat(op.forward_fp32(temb.bfloat16().float(), weight, bias)).bfloat16() + assert (ours == declared).float().mean() > 0.999 + assert (ours == early).float().mean() < 0.8 + + def test_rows_are_batch_and_position_invariant(self, block0): + """Keep each timestep modulation row identical alone and under batch reordering.""" + + weight, bias = block0 + op = _cuda_op() + temb = _temb(5, seed=3) + full = op(temb, weight, bias) + for i in range(5): + single = op(temb[i : i + 1], weight, bias) + for s, f in zip(single, full, strict=True): + assert torch.equal(s, f[3 * i : 3 * i + 3]) + swapped = op(temb.flip(0), weight, bias) + for s, f in zip(swapped, full, strict=True): + assert torch.equal(s.view(5, 3, -1).flip(0), f.view(5, 3, -1)) + + def test_backward_against_golden(self, block0): + """Check embedding and parameter gradients against a high-precision declared-cast VJP.""" + + weight, bias = block0 + temb = _temb(2, seed=5) + g = torch.Generator(device="cuda").manual_seed(0) + grads = [torch.randn(6, 5376, device="cuda", generator=g).bfloat16() for _ in range(6)] + leaves = [t.detach().clone().requires_grad_(True) for t in (temb, weight, bias)] + torch.autograd.backward(list(_cuda_op()(*leaves)), grads) + + ref = [t.detach().double().requires_grad_(True) for t in (temb, weight, bias)] + act = ref[0] * torch.sigmoid(ref[0]) + act = act + (act.to(torch.bfloat16).double() - act).detach() # declared cast, identity VJP + table = F.linear(act, ref[1], ref[2]).view(-1, 6 * 5376).chunk(6, dim=-1) + torch.autograd.backward(list(table), [gr.double() for gr in grads]) + # tolerance_contract.json gradient_accuracy / reduction: FP32 temb grad + # and BF16 weight/bias grads (rounded once from an FP32 accumulation). + torch.testing.assert_close(leaves[0].grad.double(), ref[0].grad, atol=1e-4, rtol=1e-4) + for leaf, r in zip(leaves[1:], ref[1:], strict=True): + torch.testing.assert_close(leaf.grad.double(), r.grad, atol=1e-1, rtol=2e-2) + assert (leaf.grad == r.grad.to(torch.bfloat16)).float().mean() > 0.999 + + +@requires_cuda +class TestCudaSynthetic: + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) + @pytest.mark.parametrize("hidden, dim", [(24, 16), (40, 72), (8, 2688)]) + def test_other_sizes(self, dtype, hidden, dim): + """Match the projection golden at non-checkpoint dimensions in FP32 and BF16.""" + + weight, bias = _synthetic(hidden, dim, dtype=dtype, seed=hidden) + temb = _temb(3, dim, seed=hidden) + tol = dict(atol=5e-2, rtol=2e-2) if dtype == torch.bfloat16 else dict(atol=1e-4, rtol=1e-4) + for out, gold in zip( + _cuda_op()(temb, weight, bias), + NativeH3AdaLNProjectionOp().forward_fp32(temb, weight, bias), + strict=True, + ): + torch.testing.assert_close(out.float(), gold, **tol) + + def test_repeat_forward_backward_bitwise(self): + """Require repeated projection outputs and all input gradients to match bitwise.""" + + weight, bias = _synthetic(64, 2688, seed=9) + temb = _temb(3, seed=9) + grads = [torch.randn(9, 64, device="cuda").bfloat16() for _ in range(6)] + runs = [] + for _ in range(3): + leaves = [t.detach().clone().requires_grad_(True) for t in (temb, weight, bias)] + outs = _cuda_op()(*leaves) + torch.autograd.backward(list(outs), grads) + runs.append([*[o.detach() for o in outs], *[leaf.grad for leaf in leaves]]) + for later in runs[1:]: + for a, b in zip(runs[0], later, strict=True): + assert torch.equal(a, b) + + def test_temb_grad_rows_are_batch_invariant(self): + """Keep embedding gradients bitwise equal in full batches and single-row calls.""" + + weight, bias = _synthetic(64, 2688, seed=11) + temb = _temb(4, seed=11) + grads = [torch.randn(12, 64, device="cuda").bfloat16() for _ in range(6)] + + def d_temb(rows): + """Compute selected embedding-row gradients with their matching modality gradients.""" + + leaf = temb[rows].detach().clone().requires_grad_(True) + outs = _cuda_op()(leaf, weight, bias) + torch.autograd.backward( + list(outs), [g.view(4, 3, -1)[rows].reshape(-1, 64) for g in grads] + ) + return leaf.grad + + full = d_temb(slice(0, 4)) + for i in range(4): + assert torch.equal(d_temb(slice(i, i + 1))[0], full[i]) + + @pytest.mark.parametrize("num", [8, 9, 17]) + def test_rows_invariant_across_tensor_core_row_tiles(self, num): + """The BF16 path computes rows in tiles of 8; a row's bits must not depend on its tile.""" + + weight, bias = _synthetic(64, 2688, seed=13) + temb = _temb(num, seed=num) + full = _cuda_op()(temb, weight, bias) + for i in (0, 7, num - 1): + single = _cuda_op()(temb[i : i + 1], weight, bias) + for s_out, f_out in zip(single, full, strict=True): + assert torch.equal(s_out, f_out[3 * i : 3 * i + 3]) + + def test_rejects_cpu(self): + """Reject CPU tensors at the CUDA projection boundary.""" + + weight, bias = _synthetic(8, 16, device="cpu") + with pytest.raises(ValueError): + _cuda_op()(_temb(2, 16, device="cpu"), weight, bias) + + def test_registry_dispatches_cuda(self): + """Resolve the projection registry entry to its dedicated CUDA implementation.""" + + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("adaln_projection_3mod", device="cuda") + assert type(op).__name__ == "H3AdaLNProjectionCudaOp" diff --git a/tests/models/minimax_h3/test_h3_adaln_row_gather.py b/tests/models/minimax_h3/test_h3_adaln_row_gather.py new file mode 100644 index 000000000..2b7df3803 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_adaln_row_gather.py @@ -0,0 +1,298 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``adaln_row_gather``: ``timestep_index * 3 + token_tag`` lookup of six tensors. + +* forward is a copy: bitwise equal to diffusers' six ``index_select`` calls, + for every packing, length and index dtype; +* semantic indices fail closed (out-of-range tags or timesteps, RFC probes + H2/H3), and flipping a tag or offsetting a timestep picks a different row; +* backward is a deterministic FP32 segmented sum: repeat-bitwise and + correctly rounded, unlike the atomic BF16 ``index_select`` backward; +* batch invariance: a position's output does not depend on which other + positions share the call, and a row's gradient does not depend on the + positions that reference other rows. +""" + +from __future__ import annotations + +import os +import subprocess +import sys +import textwrap +from pathlib import Path + +import pytest +import torch + +from rl_engine.reference.minimax_h3.adaln_row_gather import NativeH3AdaLNRowGatherOp +from rl_engine.validation.models.h3_cases import h3_packed_layout +from rl_engine.validation.models.h3_provider import provider_adaln_row_gather + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +def _cuda_op(): + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + H3AdaLNRowGatherCudaOp, + adaln_row_gather_available, + ) + + if not adaln_row_gather_available(): + pytest.skip("rl_engine._C lacks h3_adaln_row_gather_*") + return H3AdaLNRowGatherCudaOp() + + +def _rows(num_timesteps, hidden=5376, dtype=torch.bfloat16, seed=0, device="cuda"): + g = torch.Generator(device="cpu").manual_seed(seed) + return torch.randn(3 * num_timesteps, 6 * hidden, generator=g).to(dtype).to(device) + + +class TestReference: + def test_matches_diffusers_index_select(self): + rows = _rows(2, hidden=4, device="cpu") + ti, tags = h3_packed_layout(11, 2, device="cpu") + index = ti * 3 + tags + chunks = rows.chunk(6, dim=-1) + for ours, chunk in zip(NativeH3AdaLNRowGatherOp().forward(rows, ti, tags), chunks): + assert torch.equal(ours, chunk.index_select(0, index)) + + @pytest.mark.parametrize( + "mutate, error", + [ + (lambda ti, tags: (ti, tags.clone().fill_(3)), IndexError), # H2: no 4th modality + (lambda ti, tags: (ti, tags - 1), IndexError), + (lambda ti, tags: (ti + 2, tags), IndexError), # H3: timestep offset past T + (lambda ti, tags: (ti, tags[:-1]), ValueError), + (lambda ti, tags: (ti[:0], tags[:0]), ValueError), + (lambda ti, tags: (ti.float(), tags), TypeError), + (lambda ti, tags: (ti.int(), tags), TypeError), # mixed index dtypes + ], + ) + def test_rejects_invalid_indices(self, mutate, error): + rows = _rows(2, hidden=4, device="cpu") + ti, tags = mutate(*h3_packed_layout(9, 2, device="cpu")) + with pytest.raises(error): + NativeH3AdaLNRowGatherOp().forward(rows, ti, tags) + + def test_rejects_bad_rows(self): + ti, tags = h3_packed_layout(4, 1, device="cpu") + with pytest.raises(ValueError): + NativeH3AdaLNRowGatherOp().forward(torch.zeros(4, 24), ti, tags) # not 3 per t + with pytest.raises(ValueError): + NativeH3AdaLNRowGatherOp().forward(torch.zeros(3, 25), ti, tags) # not 6 * H + + +@requires_cuda +class TestCudaForward: + @pytest.mark.parametrize("seq", [1, 2, 3, 257, 4097, 32768]) + @pytest.mark.parametrize("num_timesteps", [1, 3]) + def test_bitwise_equal_to_index_select(self, seq, num_timesteps): + rows = _rows(num_timesteps, seed=seq) + ti, tags = h3_packed_layout(seq, num_timesteps, seed=seq) + ours = _cuda_op()(rows, ti, tags) + ref = provider_adaln_row_gather(rows.chunk(6, dim=-1), ti, tags) + assert len(ours) == 6 + for a, b in zip(ours, ref): + assert a.shape == (seq, 5376) and a.is_contiguous() and torch.equal(a, b) + + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) + @pytest.mark.parametrize("hidden", [24, 5, 5376]) + def test_dtypes_and_unaligned_hidden(self, dtype, hidden): + rows = _rows(2, hidden=hidden, dtype=dtype) + ti, tags = h3_packed_layout(33, 2) + for a, b in zip( + _cuda_op()(rows, ti.int(), tags.int()), NativeH3AdaLNRowGatherOp()(rows, ti, tags) + ): + assert torch.equal(a, b) + + def test_strided_rows_view(self): + table = _rows(4, hidden=16) + rows = table[::2] # row stride 2 * 6H, still 3 rows per timestep after slicing 12 -> 6 + ti, tags = h3_packed_layout(20, 2) + for a, b in zip(_cuda_op()(rows, ti, tags), NativeH3AdaLNRowGatherOp()(rows, ti, tags)): + assert torch.equal(a, b) + + def test_tag_flip_and_timestep_offset_select_other_rows(self): + rows = _rows(2, hidden=16) + ti, tags = h3_packed_layout(30, 2) + base = _cuda_op()(rows, ti, tags) + flipped = _cuda_op()(rows, ti, (tags + 1) % 3) # H2 + shifted = _cuda_op()(rows, 1 - ti, tags) # H3 within range + assert not torch.equal(base[0], flipped[0]) + assert not torch.equal(base[0], shifted[0]) + # Every output row is exactly the addressed table row, nothing else. + index = ti * 3 + tags + for c, out in enumerate(base): + assert torch.equal(out, rows[index, c * 16 : (c + 1) * 16]) + + def test_rejects_out_of_range_on_device(self): + rows = _rows(1, hidden=16) + ti, tags = h3_packed_layout(8, 1) + with pytest.raises(IndexError): + _cuda_op()(rows, ti, tags + 3) + + @pytest.mark.parametrize("entrypoint", ["native", "unchecked_wrapper"]) + @pytest.mark.parametrize( + "timestep, tag, index_dtype", + [ + pytest.param(-1, 0, "int32", id="negative-timestep"), + pytest.param(2, 0, "int64", id="timestep-past-table"), + pytest.param(0, 3, "int32", id="tag-past-modality-valid-row"), + pytest.param(1, -1, "int64", id="negative-tag-valid-row"), + # Multiplication by 3 wraps this value to row 2 in signed int64. + pytest.param((2**64 + 2) // 3, 0, "int64", id="timestep-overflow-valid-row"), + pytest.param(0, 2**63 - 1, "int64", id="int64-max-tag"), + ], + ) + def test_native_bounds_assertions(self, entrypoint, timestep, tag, index_dtype): + """Reject semantic index violations without poisoning pytest's CUDA context.""" + + _cuda_op() + probe = textwrap.dedent(""" + import sys + import torch + from rl_engine.backends.extension import _C + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + H3AdaLNRowGatherCudaOp, + ) + + entrypoint, timestep, tag, index_dtype = sys.argv[1:] + rows = torch.zeros((6, 6), dtype=torch.float32, device="cuda") + dtype = getattr(torch, index_dtype) + ti = torch.tensor([int(timestep)], dtype=dtype, device="cuda") + tags = torch.tensor([int(tag)], dtype=dtype, device="cuda") + if entrypoint == "native": + _C.h3_adaln_row_gather_forward(rows, ti, tags, 6, 3) + else: + H3AdaLNRowGatherCudaOp().forward(rows, ti, tags, check_range=False) + torch.cuda.synchronize() + """) + completed = subprocess.run( + [sys.executable, "-c", probe, entrypoint, str(timestep), str(tag), index_dtype], + cwd=Path(__file__).resolve().parents[3], + env={**os.environ, "CUDA_LAUNCH_BLOCKING": "1"}, + capture_output=True, + text=True, + timeout=60, + ) + assert completed.returncode != 0, "native gather accepted invalid semantic indices" + assert ( + "device-side assert triggered" in completed.stderr + ), f"expected a CUDA bounds assertion, got:\n{completed.stdout}\n{completed.stderr}" + + def test_gather_chunks_drop_in(self): + rows = _rows(2, hidden=16) + ti, tags = h3_packed_layout(10, 2) + chunks = rows.chunk(6, dim=-1) + for a, b in zip(_cuda_op().gather_chunks(chunks, ti, tags), _cuda_op()(rows, ti, tags)): + assert torch.equal(a, b) + + +@requires_cuda +class TestCudaBackward: + def _grads(self, seq, hidden, seed=0): + g = torch.Generator(device="cuda").manual_seed(seed) + return [torch.randn(seq, hidden, device="cuda", generator=g).bfloat16() for _ in range(6)] + + def _cuda_grad(self, rows, ti, tags, grads): + leaf = rows.detach().clone().requires_grad_(True) + torch.autograd.backward(list(_cuda_op()(leaf, ti, tags)), grads) + return leaf.grad + + def test_correctly_rounded_fp32_segment_sum(self): + rows = _rows(3, seed=1) + ti, tags = h3_packed_layout(4097, 3, seed=1) + grads = self._grads(4097, 5376) + ours = self._cuda_grad(rows, ti, tags, grads) + ref = rows.detach().double().requires_grad_(True) + torch.autograd.backward( + list(NativeH3AdaLNRowGatherOp().forward_fp32(ref, ti, tags)), [g.float() for g in grads] + ) + assert ours.dtype == torch.bfloat16 + # tolerance_contract.json gradient_accuracy / elementwise / bfloat16. + torch.testing.assert_close(ours.double(), ref.grad, atol=2e-2, rtol=1.6e-2) + assert (ours == ref.grad.bfloat16()).float().mean() > 0.9999 + + def test_repeat_bitwise(self): + rows = _rows(2, hidden=512, seed=2) + ti, tags = h3_packed_layout(5000, 2, seed=2) + grads = self._grads(5000, 512, seed=2) + first = self._cuda_grad(rows, ti, tags, grads) + for _ in range(3): + assert torch.equal(self._cuda_grad(rows, ti, tags, grads), first) + + def test_forward_rows_invariant_to_batch_size_and_position(self): + rows = _rows(3, hidden=512, seed=5) + ti, tags = h3_packed_layout(4097, 3, seed=5) + full = _cuda_op()(rows, ti, tags) + perm = torch.randperm(4097, generator=torch.Generator().manual_seed(5)).cuda() + for pick in ( + torch.tensor([0]), + torch.tensor([2048]), + torch.arange(100, 357), + torch.arange(4000, 4097), + perm, + ): + pick = pick.cuda() + part = _cuda_op()(rows, ti[pick], tags[pick]) + for a, b in zip(part, full): + assert torch.equal(a, b[pick]) + + def test_backward_row_gradient_independent_of_other_rows_tokens(self): + # Each row's segment is tiled on its own, so a row's gradient from the + # full packing equals the one from a packing of only its own positions + # (same order), even though ~455 positions per row cross tile boundaries. + rows = _rows(3, hidden=512, seed=6) + ti, tags = h3_packed_layout(4097, 3, seed=6) + grads = self._grads(4097, 512, seed=6) + full = self._cuda_grad(rows, ti, tags, grads) + flat = ti * 3 + tags + for r in range(rows.shape[0]): + own = (flat == r).nonzero().squeeze(1) + assert own.numel() > 256 + alone = self._cuda_grad(rows, ti[own], tags[own], [g[own] for g in grads]) + assert torch.equal(alone[r], full[r]) + assert torch.count_nonzero(alone[torch.arange(rows.shape[0], device="cuda") != r]) == 0 + + def test_unreferenced_rows_get_zero_grad(self): + rows = _rows(3, hidden=16) + ti = torch.zeros(7, dtype=torch.long, device="cuda") + tags = torch.zeros(7, dtype=torch.long, device="cuda") # only row 0 is used + grad = self._cuda_grad(rows, ti, tags, self._grads(7, 16)) + assert torch.count_nonzero(grad[1:]) == 0 + assert torch.count_nonzero(grad[0]) > 0 + + def test_segments_longer_than_one_tile(self): + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import BACKWARD_TILE + + seq = 3 * BACKWARD_TILE + 17 + rows = _rows(1, hidden=8, dtype=torch.float32) + ti = torch.zeros(seq, dtype=torch.long, device="cuda") + tags = torch.zeros(seq, dtype=torch.long, device="cuda") + grads = [torch.ones(seq, 8, device="cuda") for _ in range(6)] + grad = self._cuda_grad(rows, ti, tags, grads) + assert torch.equal(grad[0], torch.full((48,), float(seq), device="cuda")) + + def test_reference_backward_is_deterministic_and_accurate(self): + rows = _rows(2, hidden=256, seed=4) + ti, tags = h3_packed_layout(3000, 2, seed=4) + grads = self._grads(3000, 256, seed=4) + + def ref_grad(): + leaf = rows.detach().clone().requires_grad_(True) + torch.autograd.backward(list(NativeH3AdaLNRowGatherOp()(leaf, ti, tags)), grads) + return leaf.grad + + first = ref_grad() + assert torch.equal(ref_grad(), first) + torch.testing.assert_close( + first.float(), self._cuda_grad(rows, ti, tags, grads).float(), atol=1e-1, rtol=2e-2 + ) + + def test_registry_dispatches_cuda(self): + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("adaln_row_gather", device="cuda") + assert type(op).__name__ == "H3AdaLNRowGatherCudaOp" diff --git a/tests/models/minimax_h3/test_h3_benchmark_timing.py b/tests/models/minimax_h3/test_h3_benchmark_timing.py new file mode 100644 index 000000000..709c1d7e9 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_benchmark_timing.py @@ -0,0 +1,259 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CPU regressions for benchmark ordering and backward timing boundaries.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from rl_engine.validation.models import h3_report + + +@pytest.fixture +def cuda_clock(monkeypatch): + clock = SimpleNamespace(time=0.0, active=False, log=[], allocated=0, peak=0, events=0) + + class Event: + def __init__(self, *, enable_timing): + assert enable_timing + self.start = clock.events % 2 == 0 + clock.events += 1 + + def record(self): + clock.active = self.start + clock.log.append("start" if self.start else "end") + self.time = clock.time + + def synchronize(self): + assert not clock.active + + def elapsed_time(self, end): + return end.time - self.time + + def reset_peak(): + clock.log.append("reset_peak") + clock.peak = clock.allocated + + def allocated(): + clock.log.append("memory_baseline") + return clock.allocated + + monkeypatch.setattr(torch.cuda, "Event", Event) + monkeypatch.setattr(torch.cuda, "synchronize", lambda: clock.log.append("synchronize")) + monkeypatch.setattr(torch.cuda, "reset_peak_memory_stats", reset_peak) + monkeypatch.setattr(torch.cuda, "memory_allocated", allocated) + monkeypatch.setattr(torch.cuda, "max_memory_allocated", lambda: clock.peak) + return clock + + +def test_setup_is_outside_backward_events_and_memory_baseline(cuda_clock): + states = [] + + def setup(): + assert not cuda_clock.active + cuda_clock.log.append("setup") + cuda_clock.time += 50.0 + cuda_clock.allocated += 64 * 2**20 + return object() + + def backward(state): + states.append(state) + cuda_clock.log.append("backward") + cuda_clock.time += 2.0 + cuda_clock.allocated += 3 * 2**20 + cuda_clock.peak = cuda_clock.allocated + + assert h3_report.time_us(backward, warmup=1, iters=2, setup=setup) == 2000.0 + assert len({id(state) for state in states}) == 3 + assert cuda_clock.log.count("start") == 2 + for index, entry in enumerate(cuda_clock.log): + if entry == "start": + assert cuda_clock.log[index - 1] == "setup" + assert cuda_clock.log[index + 1 : index + 3] == ["backward", "end"] + + cuda_clock.log.clear() + assert h3_report.peak_mib(backward, setup=setup) == 3.0 + assert cuda_clock.log == [ + "setup", + "synchronize", + "reset_peak", + "memory_baseline", + "backward", + "synchronize", + ] + + +@pytest.mark.parametrize("keys", [("candidate", "provider"), h3_report.TIMED_KEYS]) +def test_measure_interleaves_and_records_both_orders(cuda_clock, keys): + calls = [] + case = {"op": "test", "case": "tiny", "backend": "cpu", "bytes": 4096} + for position, key in enumerate(keys, start=1): + + def run(key=key, duration=position): + calls.append(key) + cuda_clock.time += duration + + case[key] = run + + row = h3_report.measure(case, warmup=2, iters=4) + first, second = list(keys), list(reversed(keys)) + assert calls == (first + second) * 3 + first # warmup, samples, then memory calls + assert row["execution_order"] == { + "policy": "alternating", + "iteration_0": first, + "iteration_1": second, + } + for position, key in enumerate(keys, start=1): + assert row[f"{key}_us"] == position * 1000.0 + assert row[f"{key}_gbps"] == pytest.approx(4096 / (position * 1e-3) / 1e9) + assert row[f"{key}_peak_mib"] == 0.0 + assert not any(key.endswith("_setup") for key in row) + + +@pytest.mark.parametrize("operator", ["norm", "gather", "gate", "final"]) +def test_backward_perf_cases_prepare_fresh_cpu_graphs_before_events( + monkeypatch, cuda_clock, operator +): + prepared_leaves = [] + previous_threads = torch.get_num_threads() + torch.set_num_threads(1) + + def track_backward(outputs): + tensors = outputs if isinstance(outputs, tuple) else (outputs,) + + def backward_hook(grad): + cuda_clock.time += 2.0 / len(tensors) + return grad + + for tensor in tensors: + tensor.register_hook(backward_hook) + return outputs + + def track(fn, leaves, *args): + if leaves[0].requires_grad: + assert not cuda_clock.active + prepared_leaves.append([leaf if leaf.is_leaf else leaf._base for leaf in leaves]) + cuda_clock.time += 1.0 + outputs = fn(*leaves, *args) + return track_backward(outputs) if leaves[0].requires_grad else outputs + + if operator == "norm": + + def inputs(seq): + return ( + torch.randn(1, seq, 4), + torch.ones(4), + torch.randn(3, 4), + torch.randn(3, 4), + torch.arange(seq) % 3, + ) + + provider = h3_report.provider_norm_modulate + + def forward(x, weight, shift, scale, index, **kwargs): + return track(provider, [x, weight, shift, scale], index) + + monkeypatch.setattr(h3_report, "NORM_SEQ_LENS", (2, 3)) + monkeypatch.setattr(h3_report, "_norm_inputs", inputs) + monkeypatch.setattr(h3_report, "provider_norm_modulate", forward) + op = SimpleNamespace(forward_modulated=forward) + build = h3_report._norm_perf + elif operator == "gate": + + def inputs(seq): + return ( + torch.randn(3, 24), + torch.randn(1, seq, 4), + torch.randn(1, seq, 4), + torch.arange(seq) % 3, + ) + + provider = h3_report.provider_gate_residual + + def forward(residual, y, gate, index, **kwargs): + return track(lambda r, y_, g: provider(r, g, index, y_), [residual, y, gate]) + + monkeypatch.setattr(h3_report, "NORM_SEQ_LENS", (2, 3)) + monkeypatch.setattr(h3_report, "_gate_inputs", inputs) + monkeypatch.setattr( + h3_report, "provider_gate_residual", lambda r, g, i, y: forward(r, y, g, i) + ) + op = SimpleNamespace(forward=forward) + build = h3_report._gate_perf + elif operator == "final": + + def inputs(seq): + return ( + torch.randn(1, seq, 4), + torch.ones(4), + torch.randn(3, 2), + torch.randn(8, 2), + torch.randn(8), + torch.arange(seq) % 3, + ) + + def forward(x, nw, temb, w, b, ti): + def compute(x, nw, temb, w, b): + rows = torch.nn.functional.linear(torch.nn.functional.silu(temb), w, b) + shift, scale = rows.chunk(2, dim=-1) + norm = torch.nn.functional.rms_norm(x, (4,), nw) + return norm * (1 + scale[ti]) + shift[ti] + + return track(compute, [x, nw, temb, w, b]) + + monkeypatch.setattr(h3_report, "NORM_SEQ_LENS", (2, 3)) + monkeypatch.setattr(h3_report, "_final_inputs", inputs) + monkeypatch.setattr(h3_report, "provider_final_adaln_out", forward) + op = forward + build = h3_report._final_perf + else: + randn = torch.randn + + def cpu_randn(*args, **kwargs): + kwargs["device"] = "cpu" + return randn(*args, **kwargs) + + provider = h3_report.provider_adaln_row_gather + + def forward(rows, ti, tags, **kwargs): + def gather(leaf, ti, tags): + return provider(leaf.chunk(6, dim=-1), ti, tags) + + return track(gather, [rows], ti, tags) + + def provider_forward(chunks, ti, tags): + if chunks[0].requires_grad: + assert not cuda_clock.active + prepared_leaves.append([chunks[0]._base]) + cuda_clock.time += 1.0 + outputs = provider(chunks, ti, tags) + return track_backward(outputs) if chunks[0].requires_grad else outputs + + monkeypatch.setattr(torch, "randn", cpu_randn) + monkeypatch.setattr(h3_report, "GATHER_SEQ_LENS", (2, 3)) + monkeypatch.setattr( + h3_report, + "h3_packed_layout", + lambda seq, num, seed: (torch.arange(seq) % num, torch.arange(seq) % 3), + ) + monkeypatch.setattr(h3_report, "provider_adaln_row_gather", provider_forward) + op = SimpleNamespace(forward=forward) + build = h3_report._gather_perf + + try: + registry = SimpleNamespace(get_op=lambda *args, **kwargs: op) + for case in build(registry): + row = h3_report.measure(case, warmup=1, iters=2) + assert row["backward_timing_scope"] == "backward_only" + assert row["candidate_backward_us"] == pytest.approx(2000.0) + assert row["provider_backward_us"] == pytest.approx(2000.0) + finally: + torch.set_num_threads(previous_threads) + + assert prepared_leaves + assert len({id(leaves[0]) for leaves in prepared_leaves}) == len(prepared_leaves) + assert all(leaf.grad is not None for leaves in prepared_leaves for leaf in leaves) diff --git a/tests/models/minimax_h3/test_h3_chain_replay_cli.py b/tests/models/minimax_h3/test_h3_chain_replay_cli.py new file mode 100644 index 000000000..b9baea84d --- /dev/null +++ b/tests/models/minimax_h3/test_h3_chain_replay_cli.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Reject incomplete chain dependencies before checking CUDA or loading weights.""" + +from __future__ import annotations + +import sys + +import pytest + +from tools.validation.models import h3_chain_replay + +STAGE_NAMES = [stage.name for stage in h3_chain_replay.h3_chain.STAGES] + + +@pytest.fixture +def no_cuda_check(monkeypatch): + """Fail if invalid CLI arguments reach a CUDA availability check.""" + + def unexpected_check(): + """Report an environment check performed before argument rejection.""" + + pytest.fail("invalid arguments must be rejected before checking CUDA") + + monkeypatch.setattr(h3_chain_replay.torch.cuda, "is_available", unexpected_check) + + +@pytest.mark.parametrize("stages, message", [("", "unknown stages"), ("unknown", "unknown stages")]) +def test_reject_invalid_stages(monkeypatch, capsys, no_cuda_check, stages, message): + """Reject unknown stages before any device check.""" + + monkeypatch.setattr(sys, "argv", ["h3_chain_replay.py", "--stages", stages]) + with pytest.raises(SystemExit) as exc: + h3_chain_replay.main() + assert exc.value.code == 2 + assert message in capsys.readouterr().err + + +@pytest.mark.parametrize("count", range(1, len(STAGE_NAMES))) +def test_backward_requires_all_stages(monkeypatch, capsys, no_cuda_check, count): + """Require the full chain for parameter-gradient replay.""" + + monkeypatch.setattr( + sys, + "argv", + ["h3_chain_replay.py", "--stages", ",".join(STAGE_NAMES[:count]), "--backward"], + ) + with pytest.raises(SystemExit) as exc: + h3_chain_replay.main() + assert exc.value.code == 2 + assert "--backward requires every stage" in capsys.readouterr().err + + +@pytest.mark.parametrize( + "stages", + [ + *(",".join(STAGE_NAMES[:count]) for count in range(1, len(STAGE_NAMES) + 1)), + STAGE_NAMES[1], + STAGE_NAMES[-1], + ",".join(STAGE_NAMES[:1] + STAGE_NAMES[2:]), + ",".join(reversed(STAGE_NAMES)), + ",".join([STAGE_NAMES[0], STAGE_NAMES[0]]), + ], +) +def test_accept_any_known_selection(monkeypatch, stages): + """Any selection of known stages reaches the device requirement; the replay + runs its prerequisites (tests/models/minimax_h3/test_h3_cli.py).""" + + monkeypatch.setattr(h3_chain_replay.torch.cuda, "is_available", lambda: False) + monkeypatch.setattr(sys, "argv", ["h3_chain_replay.py", "--stages", stages]) + with pytest.raises(SystemExit, match="the chain replay needs a CUDA device"): + h3_chain_replay.main() + + +@pytest.mark.parametrize("options", [[], ["--backward"]]) +def test_accept_all_stages_by_default(monkeypatch, options): + """Keep the complete chain as the default for forward and backward replay.""" + + monkeypatch.setattr(h3_chain_replay.torch.cuda, "is_available", lambda: False) + monkeypatch.setattr(sys, "argv", ["h3_chain_replay.py", *options]) + with pytest.raises(SystemExit, match="the chain replay needs a CUDA device"): + h3_chain_replay.main() diff --git a/tests/models/minimax_h3/test_h3_cli.py b/tests/models/minimax_h3/test_h3_cli.py new file mode 100644 index 000000000..8e97fc6a2 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_cli.py @@ -0,0 +1,174 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""The H3 CLIs preserve chain dependencies and require pinned report weights.""" + +from __future__ import annotations + +import json +import sys +from dataclasses import replace +from types import SimpleNamespace + +import pytest +import torch + +from rl_engine.validation.models import h3_chain +from rl_engine.validation.models.h3_weights import WEIGHTS_ENV +from tools.validation.models import h3_chain_replay, h3_evidence + + +@pytest.fixture +def cpu_chain_cli(monkeypatch): + """Replace CUDA chain stages with CPU stubs that track dependencies and provider drift.""" + + executed = [] + state = {"provider_shift": 0} + + class CpuOp: + def __call__(self, value): + """Advance the synthetic chain value by one for candidate and golden replay.""" + + return value + 1 + + forward_fp32 = __call__ + + stages = [] + for index, stage in enumerate(h3_chain.STAGES): + + def stage_input(ctx, upstream, index=index): + """Use raw timesteps for stage zero and require predecessor outputs thereafter.""" + + if index == 0: + return ctx["timestep"] + assert isinstance(upstream, torch.Tensor), "a preceding stage must run first" + return upstream + + stages.append( + replace( + stage, + candidate=lambda op, ctx, up, stage_input=stage_input: op(stage_input(ctx, up)), + golden=lambda op, ctx, up, stage_input=stage_input: op(stage_input(ctx, up)), + provider=lambda ctx, up, stage_input=stage_input, index=index: ( + stage_input(ctx, up) + 1 + (state["provider_shift"] if index == 0 else 0) + ), + ) + ) + + def get_op(op_type, *, device): + """Record each dispatched stage and return its deterministic CPU stub.""" + + executed.append(op_type) + return CpuOp() + + monkeypatch.setattr(h3_chain, "STAGES", stages) + monkeypatch.setattr(h3_chain, "make_context", lambda *a, **kw: {"timestep": torch.zeros(1)}) + monkeypatch.setattr(h3_chain, "golden_op", lambda *a: CpuOp()) + monkeypatch.setattr(h3_chain_replay, "KernelRegistry", lambda: SimpleNamespace(get_op=get_op)) + monkeypatch.setattr(h3_chain_replay, "load_h3_conditioning_weights", lambda *a: {}) + monkeypatch.setattr(h3_chain, "environment", lambda: {}) + monkeypatch.setattr(h3_chain, "git_state", lambda: {}) + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.backends.cuda.matmul, "allow_tf32", False) + monkeypatch.setattr(torch.backends.cudnn, "allow_tf32", False) + return executed, state + + +@pytest.mark.parametrize("provider_shift", [0, 1]) +@pytest.mark.parametrize( + "requested, execution_count", + [ + ("timestep_mlp_fp32", 2), + ("adaln_projection_3mod", 3), + ("timestep_sinusoid_h3,adaln_projection_3mod", 3), + ("timestep_mlp_fp32,timestep_sinusoid_h3", 2), + ], +) +def test_chain_selection_runs_dependencies( + cpu_chain_cli, monkeypatch, tmp_path, capsys, requested, execution_count, provider_shift +): + """Execute prerequisite stages but report requested stages and preserve first drift.""" + + executed, state = cpu_chain_cli + state["provider_shift"] = provider_shift + out = tmp_path / "chain.json" + monkeypatch.setattr( + sys, + "argv", + [ + "h3_chain_replay.py", + "--stages", + requested, + "--timesteps", + "1", + "--seq-lens", + "3", + "--out", + str(out), + ], + ) + h3_chain_replay.main() + + report = json.loads(out.read_text()) + stage_names = [stage.name for stage in h3_chain.STAGES] + expected_report = [name for name in stage_names if name in requested.split(",")] + assert executed == stage_names[:execution_count] + assert report["executed_stages"] == executed + assert report["stages"] == expected_report + case = report["cases"][0] + assert [entry["stage"] for entry in case["stages"]] == expected_report + assert all(entry["chained_vs_golden"]["bitwise_equal"] for entry in case["stages"]) + assert case["first_drift"] == (stage_names[0] if provider_shift else None) + assert case["first_isolated_drift"] == (stage_names[0] if provider_shift else None) + summary = capsys.readouterr().out.split("; first_drift=", 1)[0] + for name in stage_names: + assert (f"{name}=" in summary) == (name in expected_report) + + +@pytest.mark.parametrize("op", ["timestep_mlp_fp32", "adaln_projection_3mod"]) +@pytest.mark.parametrize("weights_env", [None, " "]) +def test_weighted_evidence_requires_pinned_weights(monkeypatch, tmp_path, op, weights_env): + """Reject weighted evidence without a configured checkpoint before writing a report.""" + + if weights_env is None: + monkeypatch.delenv(WEIGHTS_ENV, raising=False) + else: + monkeypatch.setenv(WEIGHTS_ENV, weights_env) + out = tmp_path / "report.json" + monkeypatch.setattr(sys, "argv", ["h3_evidence.py", "--op", op, "--out", str(out)]) + with pytest.raises(SystemExit, match=WEIGHTS_ENV): + h3_evidence.main() + assert not out.exists() + + +@pytest.mark.parametrize( + "op, weight_source", + [ + ("timestep_sinusoid_h3", "not_applicable"), + ("timestep_mlp_fp32", "pinned_checkpoint"), + ("adaln_projection_3mod", "pinned_checkpoint"), + ], +) +def test_evidence_records_weight_source(monkeypatch, tmp_path, op, weight_source): + """Label weighted evidence as pinned and sinusoid evidence as weight-independent.""" + + if weight_source == "not_applicable": + monkeypatch.delenv(WEIGHTS_ENV, raising=False) + else: + monkeypatch.setenv(WEIGHTS_ENV, str(tmp_path)) + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.backends.cuda.matmul, "allow_tf32", False) + monkeypatch.setattr(torch.backends.cudnn, "allow_tf32", False) + monkeypatch.setattr(h3_evidence, "KernelRegistry", lambda: object()) + monkeypatch.setattr( + h3_evidence, + "git_state", + lambda: {"rl_kernel_commit": "test_commit", "tracked_tree_dirty": False}, + ) + monkeypatch.setattr(h3_evidence, "environment", lambda: {}) + monkeypatch.setitem(h3_evidence.ACCURACY, op, lambda registry: {}) + monkeypatch.setitem(h3_evidence.PERF_CASES, op, lambda registry: []) + out = tmp_path / "report.json" + monkeypatch.setattr(sys, "argv", ["h3_evidence.py", "--op", op, "--out", str(out)]) + h3_evidence.main() + assert json.loads(out.read_text())["weight_source"] == weight_source diff --git a/tests/models/minimax_h3/test_h3_conditioning_e2e.py b/tests/models/minimax_h3/test_h3_conditioning_e2e.py new file mode 100644 index 000000000..e26b00506 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_conditioning_e2e.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""End to end: the H3 conditioning chain on the pinned checkpoint, through the registry. + +Runs every stage registered in ``rl_engine.validation.models.h3_chain.STAGES`` (each +RFC #420 row adds its own), chained exactly as the model calls them, and +checks every stage's promises: + +* the registry dispatched the CUDA backend; +* a repeated run is bitwise equal; +* the chained output is within the stage's contract tolerance of the golden; +* stages that promise it are bitwise equal to diffusers on diffusers' inputs, + so the first isolated drift can only be a stage that declares a reduction. +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.runtime.registry import KernelRegistry +from rl_engine.validation.models import h3_chain + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 8, + reason="needs an SM80+ CUDA device", +) + +# (T distinct timesteps, S packed rows): single row, short packed, realistic-size packed. +CASES = [(1, 3), (3, 257), (4, 4097)] + + +@pytest.fixture(scope="module") +def chain_weights(h3_weights_cpu): + """Move all pinned conditioning tensors to CUDA once for the end-to-end module.""" + + return {name: tensor.cuda() for name, tensor in h3_weights_cpu.items()} + + +@pytest.fixture(scope="module") +def registry(): + """Provide the production registry used to dispatch every conditioning stage.""" + + return KernelRegistry() + + +@pytest.mark.parametrize("num_timesteps, seq_len", CASES) +def test_chain_forward(registry, chain_weights, num_timesteps, seq_len): + """Check CUDA dispatch, golden tolerance, repeatability, and declared provider parity.""" + + report = h3_chain.run_case( + registry, chain_weights, num_timesteps=num_timesteps, seq_len=seq_len + ) + stages = {stage.name: stage for stage in h3_chain.STAGES} + for entry in report["stages"]: + stage = stages[entry["stage"]] + assert entry["backend"].endswith("CudaOp"), entry + assert entry["repeat_bitwise_equal"], entry["stage"] + assert entry["chained_vs_golden"]["within_tolerance"], entry + if stage.provider_bitwise_isolated: + assert entry["isolated_vs_provider"]["bitwise_equal"], entry["stage"] + drift = report["first_isolated_drift"] + assert drift is None or not stages[drift].provider_bitwise_isolated + + +@pytest.mark.parametrize("num_timesteps, seq_len", [(1, 257), (3, 4097)]) +def test_chain_backward(registry, chain_weights, num_timesteps, seq_len): + """Parameter gradients of the whole chain: deterministic, and FP32-accurate when fused.""" + + report = h3_chain.run_backward_case( + registry, chain_weights, num_timesteps=num_timesteps, seq_len=seq_len + ) + for name, entry in report["leaves"].items(): + for mode in ("candidate", "candidate_fused"): + assert entry[mode]["repeat_bitwise_equal"], (mode, name) + fused = entry["candidate_fused"] + if name.startswith("time_embedder"): + # FP32 parameters: no BF16 rounding anywhere on the fused path. + assert fused["max_abs_vs_golden_over_absmax"] < 1e-5, (name, fused) + else: + # BF16 parameters: only their own final rounding remains. + assert fused["correctly_rounded_fraction"] > 0.99, (name, fused) diff --git a/tests/models/minimax_h3/test_h3_det_linear.py b/tests/models/minimax_h3/test_h3_det_linear.py new file mode 100644 index 000000000..f9bc79829 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_det_linear.py @@ -0,0 +1,68 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Native input validation for the H3 deterministic linear kernels.""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.backends.extension import _C +from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import det_linear_available + +pytestmark = [ + pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device"), + pytest.mark.skipif(not det_linear_available(), reason="rl_engine._C lacks h3_det_linear_*"), +] + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_backward_input_rejects_zero_rows(dtype): + """Reject empty gradient batches in native input-gradient kernels for both dtypes.""" + + grad = torch.empty(0, 16, device="cuda", dtype=torch.float32) + weight = torch.empty(16, 8, device="cuda", dtype=dtype) + with pytest.raises(RuntimeError, match="grad must have at least one row"): + _C.h3_det_linear_backward_input(grad, weight, dtype) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("with_bias", [False, True]) +def test_backward_weight_rejects_zero_rows(dtype, with_bias): + """Reject empty gradient batches before parameter-gradient or bias-gradient reduction.""" + + grad = torch.empty(0, 16, device="cuda", dtype=torch.float32) + x = torch.empty(0, 8, device="cuda", dtype=dtype) + with pytest.raises(RuntimeError, match="grad must have at least one row"): + _C.h3_det_linear_backward_weight(grad, x, dtype, with_bias) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("out_dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("k_in,n_out", [(0, 16), (8, 0), (0, 0)]) +def test_backward_input_handles_empty_columns(dtype, out_dtype, k_in, n_out): + """Match the input VJP of empty linear dimensions in the requested output dtype.""" + + grad = torch.arange(3 * n_out, device="cuda", dtype=torch.float32).reshape(3, n_out) + weight = torch.empty(n_out, k_in, device="cuda", dtype=dtype) + expected = (grad @ weight.float()).to(out_dtype) + actual = _C.h3_det_linear_backward_input(grad, weight, out_dtype) + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("w_dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("k_in,n_out", [(0, 16), (8, 0), (0, 0)]) +@pytest.mark.parametrize("with_bias", [False, True]) +def test_backward_weight_handles_empty_columns(dtype, w_dtype, k_in, n_out, with_bias): + """Return empty weight gradients while preserving any nonempty bias reduction.""" + + grad = torch.arange(3 * n_out, device="cuda", dtype=torch.float32).reshape(3, n_out) + x = torch.arange(3 * k_in, device="cuda", dtype=dtype).reshape(3, k_in) + expected_weight = (grad.T @ x.float()).to(w_dtype) + actual = _C.h3_det_linear_backward_weight(grad, x, w_dtype, with_bias) + assert len(actual) == (2 if with_bias else 1) + torch.testing.assert_close(actual[0], expected_weight, atol=0, rtol=0) + if with_bias: + torch.testing.assert_close(actual[1], grad.sum(dim=0).to(w_dtype), atol=0, rtol=0) diff --git a/tests/models/minimax_h3/test_h3_final_adaln_out.py b/tests/models/minimax_h3/test_h3_final_adaln_out.py new file mode 100644 index 000000000..5095d2c44 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_final_adaln_out.py @@ -0,0 +1,201 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``final_adaln_out``: norm_out shift/scale projection + final norm/modulation. + +* layout: ``shift`` is the first half of ``norm_out.linear``'s output, rows + are indexed by ``timestep_indices`` (not ``adaln_indices``); +* forward matches diffusers' ``norm_out`` to the projection's 1-ULP ties + (the norm and modulation are bitwise; the GEMV tree differs from cuBLAS); +* backward is deterministic, with the table gradient kept in FP32; +* rows are batch/position invariant; malformed inputs fail closed. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.reference.minimax_h3.final_adaln_out import NativeH3FinalAdaLNOutOp +from rl_engine.validation.models.h3_cases import h3_packed_layout +from rl_engine.validation.models.h3_provider import provider_final_adaln_out + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +def _cuda_op(): + from rl_engine.backends.cuda.model_specific.minimax_h3.final_adaln_out import ( + H3FinalAdaLNOutCudaOp, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import h3_rmsnorm_available + + if not h3_rmsnorm_available(): + pytest.skip("rl_engine._C lacks h3_rmsnorm_*") + return H3FinalAdaLNOutCudaOp() + + +@pytest.fixture(scope="module") +def norm_out(h3_weights_cpu): + return [ + h3_weights_cpu[name].cuda() + for name in ("norm_out.norm.weight", "norm_out.linear.weight", "norm_out.linear.bias") + ] + + +def _inputs(seq, num_timesteps=3, seed=0, batch=1): + g = torch.Generator(device="cpu").manual_seed(seed) + x = (torch.randn(batch, seq, 5376, generator=g) * 2).bfloat16().cuda() + temb = (torch.randn(num_timesteps, 2688, generator=g) * 2).cuda() + ti, _ = h3_packed_layout(seq, num_timesteps, seed=seed) + return x, temb, ti + + +def _golden(x, norm_weight, temb, weight, bias, timestep_indices): + act = temb * torch.sigmoid(temb) + act = act + (act.to(torch.bfloat16).double() - act).detach() + table = F.linear(act, weight, bias) + table = table + (table.to(torch.bfloat16).double() - table).detach() + shift, scale = table.chunk(2, dim=-1) + n = x * torch.rsqrt(x.square().mean(-1, keepdim=True) + 1e-5) * norm_weight + return n * (1 + scale.index_select(0, timestep_indices)) + shift.index_select( + 0, timestep_indices + ) + + +class TestReference: + def test_rejects_bad_inputs(self): + op = NativeH3FinalAdaLNOutOp() + x, nw = torch.zeros(1, 4, 8), torch.ones(8) + temb, w, b = torch.zeros(2, 16), torch.zeros(16, 16), torch.zeros(16) + ti = torch.tensor([0, 1, 1, 0]) + with pytest.raises(TypeError): # BF16 temb (probe H7) + op(x, nw, temb.bfloat16(), w, b, ti) + with pytest.raises(ValueError): # not 2H rows + op(x, nw, temb, w[:10], b[:10], ti) + with pytest.raises(IndexError): + op(x, nw, temb, w, b, ti + 2) + + +@requires_cuda +class TestCuda: + def test_shift_first_layout_and_timestep_indexing(self): + """Zero weights: the bias halves are (shift, scale); rows follow timestep_indices.""" + + hidden, seq = 16, 6 + x = torch.randn(1, seq, hidden, device="cuda").bfloat16() + nw = torch.ones(hidden, device="cuda", dtype=torch.bfloat16) + w = torch.zeros(2 * hidden, 32, device="cuda", dtype=torch.bfloat16) + b = torch.cat([torch.full((hidden,), 1.0), torch.full((hidden,), -1.0)]).bfloat16().cuda() + ti = torch.tensor([0, 1, 0, 1, 1, 0], device="cuda") + out = _cuda_op()(x, nw, torch.zeros(2, 32, device="cuda"), w, b, ti) + # scale = -1 -> norm(x) * 0 + shift = 1 everywhere + assert torch.equal(out, torch.ones_like(out)) + + @pytest.mark.parametrize("seq, num_timesteps", [(1, 1), (257, 2), (4097, 3)]) + def test_forward_against_diffusers_and_golden(self, norm_out, seq, num_timesteps): + x, temb, ti = _inputs(seq, num_timesteps, seed=seq) + ours = _cuda_op()(x, *norm_out[:1], temb, *norm_out[1:], ti) + theirs = provider_final_adaln_out(x, norm_out[0], temb, *norm_out[1:], ti) + golden = NativeH3FinalAdaLNOutOp().forward_fp32(x, norm_out[0], temb, *norm_out[1:], ti) + assert (ours == theirs).float().mean() > 0.999 + torch.testing.assert_close(ours.float(), golden, atol=5e-2, rtol=2e-2) + + def test_rows_are_batch_and_position_invariant(self, norm_out): + x, temb, ti = _inputs(600, seed=1, batch=2) + full = _cuda_op()(x, norm_out[0], temb, *norm_out[1:], ti) + part = _cuda_op()(x[1:2, 200:260], norm_out[0], temb, *norm_out[1:], ti[200:260]) + assert torch.equal(part[0], full[1, 200:260]) + + def test_backward_deterministic_and_fp32_table_gradient(self, norm_out): + x, temb, ti = _inputs(4097, seed=2) + grad = torch.randn_like(x) + + def grads(fn, dtype=None): + tensors = [x, norm_out[0], temb, *norm_out[1:]] + leaves = [ + (t if dtype is None else t.to(dtype)).detach().clone().requires_grad_(True) + for t in tensors + ] + fn(*leaves).backward(grad if dtype is None else grad.to(dtype)) + return [leaf.grad for leaf in leaves] + + ours = grads(lambda *t: _cuda_op()(*t, ti)) + again = grads(lambda *t: _cuda_op()(*t, ti)) + ref = grads(lambda *t: _golden(*t, ti), torch.float64) + assert all(torch.equal(a, b) for a, b in zip(ours, again)) + for name, g, r in zip(("dx", "d_norm_weight", "d_temb", "dW", "db"), ours, ref): + rel = ((g.double() - r).abs().max() / r.abs().max()).item() + assert rel < 1e-2, (name, rel) # BF16 output rounding and round(1 + scale) only + + def test_report_backward_uses_fp64_golden(self, monkeypatch): + from types import SimpleNamespace + + from rl_engine.validation.models import h3_report + + rng = torch.Generator(device="cuda").manual_seed(3) + x = torch.randn(1, 160, 8, device="cuda", generator=rng).bfloat16() + nw = torch.randn(8, device="cuda", generator=rng).bfloat16() + temb = torch.randn(3, 4, device="cuda", generator=rng) + w = torch.randn(16, 4, device="cuda", generator=rng).bfloat16() + b = torch.randn(16, device="cuda", generator=rng).bfloat16() + ti = torch.arange(160, device="cuda") % 3 + inputs = (x, nw, temb, w, b) + native = NativeH3FinalAdaLNOutOp() + registry = SimpleNamespace( + get_op=lambda *args, **kwargs: native, + _get_or_create_backend=lambda _: native, + _priority_map={"cpu": {"final_adaln_out": [None]}}, + ) + monkeypatch.setattr(h3_report, "_final_inputs", lambda *args, **kwargs: (*inputs, ti)) + monkeypatch.setattr(h3_report, "provider_final_adaln_out", native) + report = h3_report._final_accuracy(registry) + grad = torch.randn( + x.shape, device="cuda", generator=torch.Generator(device="cuda").manual_seed(7) + ).to(x.dtype) + leaves = [t.detach().clone().requires_grad_(True) for t in inputs] + native(*leaves, ti).backward(grad) + ref = [t.double().detach().requires_grad_(True) for t in inputs] + _golden(*ref, ti).backward(grad.double()) + for name, leaf, reference in zip(("dx", "d_norm_w", "d_temb", "dW", "db"), leaves, ref): + assert reference.grad.dtype == torch.float64 + expected = float( + (leaf.grad.double() - reference.grad).abs().max() / reference.grad.abs().max() + ) + for backend in ("cuda", "provider"): + assert report["backward"][backend]["rel_error"][name] == pytest.approx(expected) + + def test_registry_dispatches_cuda(self): + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("final_adaln_out", device="cuda") + assert type(op).__name__ == "H3FinalAdaLNOutCudaOp" + + @pytest.mark.parametrize("seq", [257, 1024, 4097]) + def test_bf16_gradients_meet_gtest_contract(self, seq): + # The check_operator path; S=257 used to miss d_norm_weight before the + # golden stored norm(x) and 1 + scale in BF16 like the model does. + import argparse + + from rl_engine.validation.operators import run_operator_suite + from rl_engine.validation.operators.operator_specs import make_candidate, make_operator_case + + _cuda_op() + args = argparse.Namespace( + op="final_adaln_out", + candidate="cuda", + arch_key=None, + batch=3, + seq=seq, + normalized_dim=5376, + seed=123, + input_mode="random", + ) + report = run_operator_suite( + "final_adaln_out", + candidates=[make_candidate(args)], + cases=[make_operator_case(args, torch.bfloat16, torch.device("cuda"))], + check_grad=True, + ) + assert report.passed, report.candidates[0].cases[0].outputs diff --git a/tests/models/minimax_h3/test_h3_native_validation.py b/tests/models/minimax_h3/test_h3_native_validation.py new file mode 100644 index 000000000..bcb06f1f2 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_native_validation.py @@ -0,0 +1,380 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Exercise H3 validation at the native entrypoints, bypassing Python guards.""" + +from __future__ import annotations + +import pytest +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +@pytest.fixture +def native_extension(): + from rl_engine.backends.extension import _C, _EXT_AVAILABLE + + names = ( + "h3_rmsnorm_forward", + "h3_rmsnorm_backward", + "h3_gate_residual_forward", + "h3_gate_residual_backward", + ) + if not _EXT_AVAILABLE or not all(hasattr(_C, name) for name in names): + pytest.skip("rl_engine._C lacks native H3 entrypoints") + return _C + + +@pytest.fixture +def native_case(): + device = "cuda:0" + x = torch.arange(16, device=device, dtype=torch.float32).view(2, 8) / 10 + return { + "x": x, + "weight": torch.ones(8, device=device), + "shift": torch.zeros(2, 8, device=device), + "scale": torch.zeros(2, 8, device=device), + "gate": torch.ones(2, 8, device=device), + "grad": torch.ones_like(x), + "rstd": torch.rsqrt(x.square().mean(-1) + 1e-5), + "index": torch.arange(2, device=device, dtype=torch.int64), + "sorted_pos": torch.arange(2, device=device, dtype=torch.int64), + "tile_begin": torch.arange(2, device=device, dtype=torch.int64), + "tile_end": torch.arange(1, 3, device=device, dtype=torch.int64), + "seg_first_tile": torch.arange(3, device=device, dtype=torch.int64), + } + + +def _native_call(extension, operation, boundary, case): + x, index = case["x"], case["index"] + tiles = tuple(case[name] for name in ("sorted_pos", "tile_begin", "tile_end", "seg_first_tile")) + if operation == "rmsnorm": + modulation = (case["shift"], case["scale"], index) + if boundary == "forward": + return extension.h3_rmsnorm_forward(x, case["weight"], 1e-5, *modulation) + return extension.h3_rmsnorm_backward( + case["grad"], x, case["weight"], case["rstd"], *modulation, *tiles + ) + if boundary == "forward": + return extension.h3_gate_residual_forward(x, x, case["gate"], index) + return extension.h3_gate_residual_backward(case["grad"], x, case["gate"], index, *tiles) + + +def _other_cuda(tensor): + if torch.cuda.device_count() < 2: + pytest.skip("needs two CUDA devices") + return tensor.to("cuda:1") + + +@pytest.mark.parametrize("operation", ["rmsnorm", "gate_residual"]) +@pytest.mark.parametrize("boundary", ["forward", "backward"]) +@pytest.mark.parametrize("index_case", ["empty", "cpu", "other_cuda"]) +def test_native_rejects_invalid_row_index( + native_extension, native_case, operation, boundary, index_case +): + index = native_case["index"] + if index_case == "empty": + native_case["index"] = index[:0] + elif index_case == "cpu": + native_case["index"] = index.cpu() + else: + native_case["index"] = _other_cuda(index) + with pytest.raises( + RuntimeError, match="index must be a non-empty contiguous int64 tensor on cuda:0" + ): + _native_call(native_extension, operation, boundary, native_case) + + +@pytest.mark.parametrize("operation", ["rmsnorm", "gate_residual"]) +@pytest.mark.parametrize("metadata", ["sorted_pos", "tile_begin", "tile_end", "seg_first_tile"]) +@pytest.mark.parametrize("device", ["cpu", "other_cuda"]) +def test_native_rejects_mixed_device_tiles( + native_extension, native_case, operation, metadata, device +): + tensor = native_case[metadata] + native_case[metadata] = tensor.cpu() if device == "cpu" else _other_cuda(tensor) + with pytest.raises( + RuntimeError, match="tile metadata must be contiguous int64 tensors on cuda:0" + ): + _native_call(native_extension, operation, "backward", native_case) + + +@pytest.mark.parametrize("operation", ["rmsnorm", "gate_residual"]) +def test_native_accepts_same_device_inputs(native_extension, native_case, operation): + x = native_case["x"].detach().clone().requires_grad_(True) + if operation == "rmsnorm": + weight = native_case["weight"].detach().clone().requires_grad_(True) + shift = native_case["shift"].detach().clone().requires_grad_(True) + scale = native_case["scale"].detach().clone().requires_grad_(True) + norm = torch.nn.functional.rms_norm(x, (8,), weight, 1e-5) + ref = norm * (1 + scale) + shift + ref.backward(native_case["grad"]) + expected_grads = (x.grad, weight.grad, shift.grad, scale.grad) + got, _ = _native_call(native_extension, operation, "forward", native_case) + else: + gate = native_case["gate"].detach().clone().requires_grad_(True) + ref = native_case["x"] + gate * x + ref.backward(native_case["grad"]) + expected_grads = (x.grad, gate.grad) + got = _native_call(native_extension, operation, "forward", native_case) + torch.testing.assert_close(got, ref) + actual_grads = _native_call(native_extension, operation, "backward", native_case) + for actual, expected in zip(actual_grads, expected_grads, strict=True): + torch.testing.assert_close(actual, expected) + + +TILE_NAMES = ("sorted_pos", "tile_begin", "tile_end", "seg_first_tile") + + +@pytest.fixture +def native(): + from rl_engine.backends.extension import _C, _EXT_AVAILABLE + + if not _EXT_AVAILABLE or not all( + hasattr(_C, name) for name in ("h3_rmsnorm_forward", "h3_rmsnorm_backward") + ): + pytest.skip("rl_engine._C lacks h3_rmsnorm_*") + return _C + + +@pytest.fixture +def inputs(native): + x = torch.randn(4, 8, device="cuda") + weight = torch.ones(8, device=x.device) + _, rstd = native.h3_rmsnorm_forward(x, weight, 1e-5) + shift, scale = torch.zeros(2, 16, device=x.device).chunk(2, dim=1) + return { + "grad": torch.ones_like(x), + "x": x, + "weight": weight, + "rstd": rstd, + "shift": shift, + "scale": scale, + "index": torch.tensor([0, 1], dtype=torch.int64, device=x.device), + "sorted_pos": torch.tensor([0, 2, 1, 3], dtype=torch.int64, device=x.device), + "tile_begin": torch.tensor([0, 2], dtype=torch.int64, device=x.device), + "tile_end": torch.tensor([2, 4], dtype=torch.int64, device=x.device), + "seg_first_tile": torch.tensor([0, 1, 2], dtype=torch.int64, device=x.device), + } + + +def _call(native, inputs, entrypoint): + if entrypoint == "forward": + return native.h3_rmsnorm_forward( + inputs["x"], + inputs["weight"], + 1e-5, + inputs["shift"], + inputs["scale"], + inputs["index"], + ) + return native.h3_rmsnorm_backward(**inputs) + + +def _malformed(tensor, defect): + if defect == "cpu": + return tensor.cpu() + if defect == "dtype": + dtype = torch.float16 if tensor.dtype == torch.float32 else torch.int32 + return tensor.to(dtype) + if defect == "rank": + return tensor.unsqueeze(0) + if defect == "stride": + return torch.stack((tensor, tensor), dim=1)[:, 0] + if defect == "length": + return tensor[:-1] + raise ValueError(f"unknown defect: {defect}") + + +@pytest.mark.cuda_only +def test_valid_plain_and_modulated_backward(native, inputs): + plain = {name: inputs[name] for name in ("grad", "x", "weight", "rstd")} + assert len(native.h3_rmsnorm_backward(**plain)) == 2 + dx, dweight, dshift, dscale = native.h3_rmsnorm_backward(**inputs) + assert dx.shape == inputs["x"].shape + assert dweight.shape == inputs["weight"].shape + assert dscale.shape == inputs["scale"].shape + torch.testing.assert_close(dshift, torch.full_like(inputs["shift"], 2.0)) + torch.cuda.synchronize(inputs["x"].device) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("entrypoint", ["forward", "backward"]) +def test_rejects_empty_row_index(native, inputs, entrypoint): + inputs["index"] = inputs["index"][:0] + with pytest.raises(RuntimeError, match="row index must be a non-empty"): + _call(native, inputs, entrypoint) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("entrypoint", ["forward", "backward"]) +def test_rejects_zero_hidden_size(native, inputs, entrypoint): + inputs["x"] = inputs["x"][:, :0].contiguous() + inputs["weight"] = inputs["weight"][:0] + with pytest.raises(RuntimeError, match="at least one column"): + _call(native, inputs, entrypoint) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("defect", ["cpu", "dtype", "rank", "stride", "length"]) +@pytest.mark.parametrize("mode", ["plain", "modulated"]) +def test_rejects_invalid_rstd(native, inputs, defect, mode): + inputs["rstd"] = _malformed(inputs["rstd"], defect) + if mode == "plain": + inputs = {name: inputs[name] for name in ("grad", "x", "weight", "rstd")} + with pytest.raises(RuntimeError, match="rstd must be contiguous float32"): + native.h3_rmsnorm_backward(**inputs) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("name", TILE_NAMES) +@pytest.mark.parametrize("defect", ["cpu", "dtype", "rank", "stride", "length"]) +def test_rejects_invalid_segment_tiles(native, inputs, name, defect): + inputs[name] = _malformed(inputs[name], defect) + message = "tile_begin and tile_end" if defect == "length" and name.startswith("tile_") else name + with pytest.raises(RuntimeError, match=message): + native.h3_rmsnorm_backward(**inputs) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("name", TILE_NAMES) +def test_rejects_missing_segment_tiles(native, inputs, name): + inputs[name] = None + with pytest.raises(RuntimeError, match="needs the sorted segment tiles"): + native.h3_rmsnorm_backward(**inputs) + + +@pytest.mark.cuda_only +@pytest.mark.parametrize("name", TILE_NAMES) +def test_plain_backward_rejects_segment_tiles(native, inputs, name): + plain = {key: inputs[key] for key in ("grad", "x", "weight", "rstd", name)} + with pytest.raises(RuntimeError, match="sorted segment tiles require modulation"): + native.h3_rmsnorm_backward(**plain) + + +@pytest.mark.cuda_only +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="needs two CUDA devices") +@pytest.mark.parametrize( + ("entrypoint", "name"), + [ + ("forward", "weight"), + ("forward", "index"), + ("backward", "weight"), + ("backward", "index"), + ("backward", "rstd"), + ("backward", "sorted_pos"), + ("backward", "tile_begin"), + ("backward", "tile_end"), + ("backward", "seg_first_tile"), + ], +) +def test_rejects_tensor_on_another_cuda_device(native, inputs, entrypoint, name): + other_device = (inputs["x"].device.index + 1) % torch.cuda.device_count() + inputs[name] = inputs[name].to(f"cuda:{other_device}") + message = "row index" if name == "index" else name + with pytest.raises(RuntimeError, match=message): + _call(native, inputs, entrypoint) + + +@pytest.mark.parametrize( + "operation", ["gate_forward", "gate_dy", "norm_forward", "norm_dx", "norm_partials"] +) +@pytest.mark.parametrize("bad_index", [-1, 2]) +def test_native_table_bounds_fail_in_isolated_process(native_extension, operation, bad_index): + import subprocess + import sys + import textwrap + + script = textwrap.dedent(""" + import sys + import torch + from rl_engine import _C + op, bad = sys.argv[1], int(sys.argv[2]) + x = torch.ones(2, 8, device="cuda") + w = torch.ones(8, device="cuda") + table = torch.ones(2, 8, device="cuda") + index = torch.tensor([bad, 0], device="cuda") + rstd = torch.ones(2, device="cuda") + rows = torch.arange(2, device="cuda") + begin = torch.tensor([0], device="cuda") + end = torch.tensor([2], device="cuda") + try: + if op == "gate_forward": + _C.h3_gate_residual_forward(x, x, table, index) + elif op == "gate_dy": + _C.h3_gate_residual_backward_dy(x, table, index) + elif op == "norm_forward": + _C.h3_rmsnorm_forward(x, w, 1e-5, table, table, index) + elif op == "norm_dx": + _C.h3_rmsnorm_backward_dx(x, x, w, rstd, table, table, index) + else: + _C.h3_rmsnorm_backward_partials( + x, x, w, rstd, table, table, index, rows, begin, end, rows, begin, end) + torch.cuda.synchronize() + except RuntimeError as exc: + if "device-side assert" in str(exc): + print("BOUNDS_ASSERT_CONFIRMED") + sys.exit(0) + raise + raise AssertionError("invalid index reached native table lookup without an assertion") + """) + result = subprocess.run( + [sys.executable, "-c", script, operation, str(bad_index)], + capture_output=True, + text=True, + timeout=60, + ) + assert result.returncode == 0, result.stdout + result.stderr + assert "BOUNDS_ASSERT_CONFIRMED" in result.stdout + + +@pytest.mark.parametrize("metadata", ["sorted_pos", "tile_begin", "tile_end", "seg_first_tile"]) +def test_row_gather_rejects_metadata_on_other_device(native_extension, metadata): + if torch.cuda.device_count() < 2: + pytest.skip("needs two CUDA devices") + grad = torch.ones(6, 2, 8, device="cuda:0") + tensors = dict( + sorted_pos=torch.tensor([0, 1], device="cuda:0"), + tile_begin=torch.tensor([0], device="cuda:0"), + tile_end=torch.tensor([2], device="cuda:0"), + seg_first_tile=torch.tensor([0, 1], device="cuda:0"), + ) + tensors[metadata] = tensors[metadata].to("cuda:1") + with pytest.raises(RuntimeError, match="grad device"): + native_extension.h3_adaln_row_gather_backward(grad, **tensors, out_dtype=torch.float32) + + +@pytest.mark.parametrize("metadata", ["sorted_pos", "tile_begin", "tile_end", "seg_first_tile"]) +@pytest.mark.parametrize("bound", ["negative", "past_end"]) +def test_row_gather_metadata_bounds_in_isolated_process(native_extension, metadata, bound): + import subprocess + import sys + import textwrap + + script = textwrap.dedent(""" + import sys + import torch + from rl_engine import _C + name, bound = sys.argv[1:] + grad = torch.ones(6, 2, 8, device="cuda") + metadata = dict(sorted_pos=torch.tensor([0, 1], device="cuda"), + tile_begin=torch.tensor([0], device="cuda"), + tile_end=torch.tensor([2], device="cuda"), + seg_first_tile=torch.tensor([0, 1], device="cuda")) + metadata[name][0] = -1 if bound == "negative" else 3 + try: + _C.h3_adaln_row_gather_backward(grad, **metadata, out_dtype=torch.float32) + torch.cuda.synchronize() + except RuntimeError as exc: + if "device-side assert" in str(exc): + print("BOUNDS_ASSERT_CONFIRMED") + sys.exit(0) + raise + raise AssertionError("invalid metadata was not rejected") + """) + result = subprocess.run( + [sys.executable, "-c", script, metadata, bound], capture_output=True, text=True, timeout=60 + ) + assert result.returncode == 0, result.stdout + result.stderr + assert "BOUNDS_ASSERT_CONFIRMED" in result.stdout diff --git a/tests/models/minimax_h3/test_h3_one_block.py b/tests/models/minimax_h3/test_h3_one_block.py new file mode 100644 index 000000000..4d16d3ac7 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_one_block.py @@ -0,0 +1,126 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""One real MiniMax-H3 block, forward and backward (RFC #420 ``ws1_one_h3_block``). + +GPU tests run block 0 of the pinned checkpoint on packed FL2VA layouts and +check the block-level promises of ``rl_engine.validation.models.h3_block``: + +* every node is bitwise repeatable, and batch row 0 of a two-row batch has the + bytes of the single-row run (forward, and the gradient reaching every node); +* nodes that promise it are bitwise equal to diffusers on diffusers' inputs; +* every node, the block output and every gradient are about as close to the + FP32 golden as diffusers itself. + +CPU tests pin the packed layouts to the pinned diffusers pipeline and check +the graph's structure. +""" + +from __future__ import annotations + +import hashlib + +import pytest +import torch + +from rl_engine.reference.minimax_h3 import block as gold +from rl_engine.validation.models import h3_block +from rl_engine.validation.models import h3_provider as prov +from rl_engine.validation.models.h3_cases import H3_BLOCK_LAYOUTS, h3_fl2va_layout + +needs_gpu = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 8, + reason="needs an SM80+ CUDA device", +) + +# sha256(position_ids float64 bytes + token_tags int64 bytes) of diffusers@f53d552's +# ``MiniMaxH3PrepareLayoutStep.build_packed_sequence`` for each layout (text tags 1, patch +# (1, 2, 2), two audio channels, audio tag 2, video tag 0, no keyframe anchors). +PINNED_LAYOUTS = { + "tiny": "5329a0a13a9996c4fea7a0e676b65aa24302d65bea2e2248ba68d12184cc7c37", + "small": "64e4e79d8fb91d5a68ae7d15bb29049e08add38da2639030e089c6ef80968931", + "medium": "16e448800695eedfb85e0aacbca9bd5415ed53fbbe25f701bcb60df5f4754a31", +} + +# Accuracy: the candidate's error against the FP32 golden may exceed diffusers' by this +# factor plus an absolute floor (relative to the golden's max magnitude). +FORWARD_FACTOR, FORWARD_FLOOR = 1.25, 1e-3 +BACKWARD_FACTOR, BACKWARD_FLOOR = 1.5, 5e-3 + + +@pytest.mark.parametrize("layout", sorted(PINNED_LAYOUTS)) +def test_layout_matches_pinned_pipeline(layout): + packed = h3_fl2va_layout(**H3_BLOCK_LAYOUTS[layout], device="cpu") + data = packed["position_ids"].numpy().tobytes() + packed["token_tags"].numpy().tobytes() + assert hashlib.sha256(data).hexdigest() == PINNED_LAYOUTS[layout] + + +def test_graph_is_well_formed(): + seen = {"hidden", "temb"} + for node in h3_block.NODES: + assert set(node.inputs) <= seen, node.name + assert node.name not in seen, node.name + assert node.binding.status in ("row", "interim", "reference"), node.name + seen.add(node.name) + assert h3_block.OUTPUT == "residual_mlp" + rows = {n.binding.rfc_row for n in h3_block.NODES if n.binding.status == "row"} + assert rows == {"adaln_projection_3mod", "h3_rmsnorm", "adaln_gate_residual"} + + +def test_partial_rope_rotates_96_of_128_channels(): + packed = h3_fl2va_layout(**H3_BLOCK_LAYOUTS["tiny"], device="cpu") + cos, sin = prov.provider_rope(packed["position_ids"]) + assert cos.shape == (packed["position_ids"].shape[0], 96) and cos.dtype == torch.float32 + gold_cos, gold_sin = gold.golden_rope_tables(packed["position_ids"]) + assert torch.equal(cos, gold_cos) and torch.equal(sin, gold_sin) + x = torch.randn(1, cos.shape[0], 2, 128).bfloat16() + for out in (prov.provider_apply_rotary(x, cos, sin), gold.golden_apply_rotary(x, cos, sin)): + assert torch.equal(out[..., 96:].float(), x[..., 96:].float()) + assert not torch.equal(out[..., :96].float(), x[..., :96].float()) + + +@pytest.fixture(scope="module") +def ops(): + return h3_block.CandidateOps() + + +@needs_gpu +def test_stack_rows_dispatch_to_cuda(ops): + backends = ops.backends() + for name in ("projection", "norm", "gate"): + assert backends[name].endswith("CudaOp"), backends + + +@needs_gpu +@pytest.mark.parametrize("layout", ["tiny", "small"]) +def test_block_forward(h3_block_weights_cuda, ops, layout): + report = h3_block.run_block_case(h3_block_weights_cuda, layout=layout, ops=ops) + nodes = {node.name: node for node in h3_block.NODES} + for entry in report["nodes"]: + name = entry["node"] + assert entry["repeat_bitwise_equal"], name + assert entry["batch_row_bitwise_equal"], name + if nodes[name].provider_bitwise_isolated: + assert entry["isolated_vs_provider"]["bitwise_equal"], name + bound = FORWARD_FACTOR * entry["provider_rel_err_vs_golden"] + FORWARD_FLOOR + assert entry["candidate_rel_err_vs_golden"] <= bound, entry + assert report["first_repeat_drift"] is None + assert report["first_batch_drift"] is None + drift = report["first_isolated_drift"] + assert drift is None or not nodes[drift].provider_bitwise_isolated + out = report["output"] + assert out["candidate_rel_err_vs_golden"] <= ( + FORWARD_FACTOR * out["provider_rel_err_vs_golden"] + FORWARD_FLOOR + ) + + +@needs_gpu +def test_block_backward(h3_block_weights_cuda, ops): + report = h3_block.run_block_backward(h3_block_weights_cuda, layout="tiny", ops=ops) + for name, entry in report["leaves"].items(): + assert entry["repeat_bitwise_equal"], name + bound = BACKWARD_FACTOR * entry["provider_rel_err_vs_golden"] + BACKWARD_FLOOR + assert entry["candidate_rel_err_vs_golden"] <= bound, (name, entry) + assert report["leaves"]["hidden"]["batch_row_bitwise_equal"] + assert report["first_grad_repeat_drift"] is None + assert report["first_grad_batch_drift"] is None diff --git a/tests/models/minimax_h3/test_h3_report.py b/tests/models/minimax_h3/test_h3_report.py new file mode 100644 index 000000000..ed74d1901 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_report.py @@ -0,0 +1,259 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Optional checkpoint handling preserves standalone operator measurements, and +comparative timing balances backend order and preserves its raw evidence.""" + +from __future__ import annotations + +import importlib.util +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.validation.models import h3_chain, h3_report +from rl_engine.validation.models.h3_cases import h3_packed_layout +from rl_engine.validation.models.h3_provider import ( + provider_adaln_row_gather, + provider_norm_modulate, +) +from rl_engine.validation.models.h3_weights import WEIGHTS_ENV + + +@pytest.fixture +def small_cpu_report(monkeypatch): + """Exercise actual report/autograd logic with small CPU tensors and provider ops.""" + + def dimension(size): + return {5376: 4, 6 * 5376: 24, 257: 160, 777: 160, 4097: 160, 32768: 160}.get(size, size) + + class CpuTorch: + nn = SimpleNamespace( + functional=SimpleNamespace( + rms_norm=lambda x, _shape, weight, eps: F.rms_norm(x, (x.shape[-1],), weight, eps) + ) + ) + + def __getattr__(self, name): + return getattr(torch, name) + + def Generator(self, device): + return torch.Generator(device="cpu") + + def randn(self, *shape, **kwargs): + if len(shape) == 1 and isinstance(shape[0], (tuple, list, torch.Size)): + shape = shape[0] + kwargs["device"] = "cpu" + return torch.randn(*(dimension(size) for size in shape), **kwargs) + + class Gather: + def __call__(self, rows, ti, tags): + return provider_adaln_row_gather(rows.chunk(6, dim=-1), ti, tags) + + forward_fp32 = __call__ + + class Norm: + def __call__(self, x, weight): + return F.rms_norm(x, (x.shape[-1],), weight, 1e-5) + + forward_modulated = staticmethod(provider_norm_modulate) + + gather = Gather() + ops = {"adaln_row_gather": gather, "h3_rmsnorm": Norm()} + registry = SimpleNamespace( + get_op=lambda name, device: ops[name], + _priority_map={"cpu": {"adaln_row_gather": ["gather"]}}, + _get_or_create_backend=lambda _backend: gather, + ) + monkeypatch.setattr(h3_report, "torch", CpuTorch()) + monkeypatch.setattr( + h3_report, + "h3_packed_layout", + lambda seq, num, seed: h3_packed_layout(dimension(seq), num, seed=seed, device="cpu"), + ) + + def norm_inputs(seq, seed=0): + generator = torch.Generator().manual_seed(seed) + weight = h3_report.h3_params(["transformer_blocks.0.norm1.weight"], [(5376,)])[0] + shift, scale = [torch.randn(9, 4, generator=generator).bfloat16() for _ in range(2)] + x = torch.randn(1, dimension(seq), 4, generator=generator).bfloat16() + ti, tags = h3_packed_layout(dimension(seq), 3, seed=seed, device="cpu") + return x, weight.bfloat16(), shift, scale, ti * 3 + tags + + monkeypatch.setattr(h3_report, "_norm_inputs", norm_inputs) + return registry + + +def test_gather_without_weights_preserves_op_results(small_cpu_report, monkeypatch): + monkeypatch.delenv(WEIGHTS_ENV, raising=False) + + def unexpected_load(*args, **kwargs): + pytest.fail("an unset checkpoint must not be loaded") + + monkeypatch.setattr(h3_report, "load_h3_conditioning_weights", unexpected_load) + result = h3_report._gather_accuracy(small_cpu_report) + assert all(result["forward_bitwise_vs_index_select"].values()) + assert set(result["op_backward"]) == {"cuda", "provider"} + assert all(row["repeat_bitwise_equal"] for row in result["op_backward"].values()) + assert result["chain_backward"] == [] + assert WEIGHTS_ENV in result["chain_backward_skipped"] + + +@pytest.mark.parametrize("configured", [False, True]) +def test_norm_reports_weight_source(small_cpu_report, monkeypatch, tmp_path, configured): + if configured: + monkeypatch.setenv(WEIGHTS_ENV, str(tmp_path)) + else: + monkeypatch.delenv(WEIGHTS_ENV, raising=False) + calls = [] + + def load_weights(device, names): + assert configured, "an unset checkpoint must not be loaded" + calls.append(names) + return {name: torch.ones(4, dtype=torch.bfloat16) for name in names} + + monkeypatch.setattr(h3_report, "load_h3_conditioning_weights", load_weights) + result = h3_report._norm_accuracy(small_cpu_report) + assert result["weight_source"] == ("pinned_checkpoint" if configured else "synthetic") + assert len(result["plain_bitwise_vs_nn_rmsnorm"]) == 4 + assert all(result["plain_bitwise_vs_nn_rmsnorm"].values()) + assert result["modulated_bitwise_vs_diffusers"] + assert result["rows_batch_invariant"] + assert set(result["backward"]) == {"cuda", "provider"} + assert bool(calls) is configured + + +def test_gather_with_weights_runs_chain(small_cpu_report, monkeypatch, tmp_path): + monkeypatch.setenv(WEIGHTS_ENV, str(tmp_path)) + weights = {"weight": torch.ones(1)} + monkeypatch.setattr(h3_report, "load_h3_conditioning_weights", lambda device: weights) + + def run_chain(registry, loaded, **case): + assert registry is small_cpu_report + assert loaded is weights + return case + + monkeypatch.setattr(h3_chain, "run_backward_case", run_chain) + result = h3_report._gather_accuracy(small_cpu_report) + assert len(result["chain_backward"]) == 5 + assert result["chain_backward"][-1] == {"num_timesteps": 4, "seq_len": 32768} + assert "chain_backward_skipped" not in result + + +@pytest.mark.parametrize("accuracy", [h3_report._gather_accuracy, h3_report._norm_accuracy]) +def test_configured_missing_weights_raise(small_cpu_report, monkeypatch, tmp_path, accuracy): + monkeypatch.setenv(WEIGHTS_ENV, str(tmp_path)) + with pytest.raises(FileNotFoundError, match="missing; run tools/weights/prepare_h3_weights.py"): + accuracy(small_cpu_report) + + +def test_plot_gather_without_chain(small_cpu_report, monkeypatch, tmp_path): + pytest.importorskip("matplotlib") + monkeypatch.delenv(WEIGHTS_ENV, raising=False) + script = ( + Path(__file__).resolve().parents[3] + / "tools" + / "validation" + / "models" + / "plot_h3_evidence.py" + ) + spec = importlib.util.spec_from_file_location("plot_h3_evidence_test", script) + plot = importlib.util.module_from_spec(spec) + spec.loader.exec_module(plot) + accuracy = h3_report._gather_accuracy(small_cpu_report) + report = { + "environment": {"gpu": "CPU fixture", "torch": torch.__version__, "cuda": "none"}, + "rl_kernel_commit": "testcommit", + "model_revision": "pinned_revision", + "accuracy": accuracy, + "perf": [ + { + "case": "S=160", + "candidate_us": 1, + "provider_us": 2, + "candidate_backward_us": 3, + "provider_backward_us": 4, + } + ], + } + figure = plot.plot_gather(report) + output = tmp_path / "figure.png" + figure.savefig(output) + assert output.stat().st_size > 0 + assert figure.axes[1].texts[0].get_text() == accuracy["chain_backward_skipped"] + plot.plt.close(figure) + + +def test_measure_interleaves_backends_and_records_samples(monkeypatch): + """Balance candidate/provider timing order and preserve samples and derived bandwidth.""" + + elapsed_us = 0.0 + calls = [] + + class Event: + def __init__(self, *, enable_timing): + """Require timing-enabled events in the deterministic timing stub.""" + + assert enable_timing + + def record(self): + """Capture the synthetic elapsed clock at this event.""" + + self.timestamp = elapsed_us + + def synchronize(self): + """Leave synchronization inert because the timing stub has no pending GPU work.""" + + pass + + def elapsed_time(self, end): + """Return elapsed synthetic event time in CUDA milliseconds.""" + + return (end.timestamp - self.timestamp) / 1e3 + + def candidate(): + """Record a candidate call and advance the synthetic clock by two microseconds.""" + + nonlocal elapsed_us + calls.append("candidate") + elapsed_us += 2.0 + + def provider(): + """Record a provider call and advance the synthetic clock by four microseconds.""" + + nonlocal elapsed_us + calls.append("provider") + elapsed_us += 4.0 + + monkeypatch.setattr(h3_report.torch.cuda, "Event", Event) + monkeypatch.setattr(h3_report.torch.cuda, "synchronize", lambda: None) + monkeypatch.setattr(h3_report, "peak_mib", lambda fn, **kwargs: 0.0) + report = h3_report.measure( + {"bytes": 1000, "candidate": candidate, "provider": provider}, warmup=1, iters=4 + ) + assert calls == [ + "candidate", + "provider", # warmup + "candidate", + "provider", + "provider", + "candidate", + "candidate", + "provider", + "provider", + "candidate", + ] + assert report["timing_order"] == [ + ["candidate", "provider"], + ["provider", "candidate"], + ["candidate", "provider"], + ["provider", "candidate"], + ] + assert report["timing_samples_us"] == {"candidate": [2.0] * 4, "provider": [4.0] * 4} + assert report["candidate_us"] == pytest.approx(2.0) + assert report["provider_us"] == pytest.approx(4.0) + assert report["candidate_gbps"] == pytest.approx(0.5) + assert report["provider_gbps"] == pytest.approx(0.25) diff --git a/tests/models/minimax_h3/test_h3_rmsnorm.py b/tests/models/minimax_h3/test_h3_rmsnorm.py new file mode 100644 index 000000000..ced9a9377 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_rmsnorm.py @@ -0,0 +1,302 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``h3_rmsnorm``: RMSNorm (eps 1e-5) and its fused AdaLN modulation. + +* forward is bitwise equal to ``nn.RMSNorm`` and to diffusers' + ``norm(x) * (1.0 + scale[i]) + shift[i]`` (block: ``adaln_indices``; + ``norm_out``: ``timestep_indices``) on the pinned norm weights; +* rows are independent of batch size and position; +* backward is deterministic and checked against an FP64 golden; +* malformed inputs fail closed. +""" + +from __future__ import annotations + +import pytest +import torch + +from rl_engine.reference.minimax_h3.rmsnorm import NativeH3RMSNormOp +from rl_engine.validation.models.h3_cases import h3_packed_layout +from rl_engine.validation.models.h3_provider import provider_norm_modulate + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +NORMS = ( + "transformer_blocks.0.norm1.weight", + "transformer_blocks.0.norm2.weight", + "token_refiner.final_norm.weight", + "norm_out.norm.weight", +) + + +def _cuda_op(): + from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import ( + H3RMSNormCudaOp, + h3_rmsnorm_available, + ) + + if not h3_rmsnorm_available(): + pytest.skip("rl_engine._C lacks h3_rmsnorm_*") + return H3RMSNormCudaOp() + + +def _x(shape, dtype=torch.bfloat16, seed=0, scale=2.0, device="cuda"): + g = torch.Generator(device="cpu").manual_seed(seed) + return (torch.randn(shape, generator=g) * scale).to(device=device, dtype=dtype) + + +def _table(rows, hidden, chunks, dtype=torch.bfloat16, seed=0): + g = torch.Generator(device="cpu").manual_seed(seed) + table = (torch.randn(rows, chunks * hidden, generator=g) * 0.5).to(dtype).cuda() + return table.chunk(chunks, dim=-1) # strided (R, H) views, like the AdaLN outputs + + +def _fp64_modulated(x, weight, shift, scale, index, eps=1e-5): + x64 = x.double() + n = x64 * torch.rsqrt(x64.square().mean(-1, keepdim=True) + eps) * weight.double() + return n * (1 + scale.double().index_select(0, index)) + shift.double().index_select(0, index) + + +class TestReference: + def test_golden_close_to_provider(self): + x, w = _x((3, 64), torch.float32, device="cpu"), torch.rand(64) + 0.5 + op = NativeH3RMSNormOp() + torch.testing.assert_close(op.forward(x, w), op.forward_fp32(x, w), atol=1e-6, rtol=1e-6) + + def test_modulated_with_prevalidated_indices(self): + x, w = _x((2, 8), torch.float32, device="cpu"), torch.ones(8) + shift, scale = torch.randn(3, 8), torch.randn(3, 8) + index = torch.tensor([2, 0]) + op = NativeH3RMSNormOp() + expected = op.forward(x, w) * (1 + scale[index]) + shift[index] + torch.testing.assert_close( + op.forward_modulated(x, w, shift, scale, index, check_range=False), + expected, + ) + + def test_rejects_bad_inputs(self): + op = NativeH3RMSNormOp() + x, w = torch.randn(2, 8), torch.ones(8) + with pytest.raises(TypeError): + op.forward(x, w.double()) + with pytest.raises(TypeError): + op.forward(x, w.bfloat16()) + with pytest.raises(ValueError): + op.forward(x, torch.ones(7)) + with pytest.raises(ValueError): + op.forward(x, w, eps=0.0) + shift, scale = torch.zeros(3, 8), torch.zeros(3, 8) + with pytest.raises(IndexError): + op.forward_modulated(x, w, shift, scale, torch.tensor([0, 3])) + with pytest.raises(ValueError): # one index per position S + op.forward_modulated(x, w, shift, scale, torch.tensor([0])) + + +@requires_cuda +class TestCudaForward: + @pytest.mark.parametrize("name", NORMS) + def test_bitwise_equal_to_nn_rmsnorm_on_pinned_weights(self, h3_weights_cpu, name): + weight = h3_weights_cpu[name].cuda() + module = torch.nn.RMSNorm(5376, eps=1e-5).cuda().bfloat16() + module.weight.data.copy_(weight) + x = _x((2, 777, 5376), seed=1) + assert torch.equal(_cuda_op()(x, weight), module(x)) + + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) + @pytest.mark.parametrize("hidden", [12, 256, 4096, 5376]) + def test_bitwise_equal_to_f_rms_norm(self, dtype, hidden): + x = _x((5, 33, hidden), dtype, seed=hidden, scale=30.0) + w = (torch.rand(hidden, device="cuda") * 2).to(dtype) + ref = torch.nn.functional.rms_norm(x, (hidden,), w, 1e-5) + assert torch.equal(_cuda_op()(x, w), ref) + + def test_block_modulation_bitwise_equal_to_diffusers(self, h3_weights_cpu): + weight = h3_weights_cpu["transformer_blocks.0.norm1.weight"].cuda() + shift, scale, *_ = _table(9, 5376, 6, seed=2) # (3T, H) rows, T = 3 + ti, tags = h3_packed_layout(4097, 3, seed=2) + index = ti * 3 + tags + x = _x((2, 4097, 5376), seed=2) + ours = _cuda_op().forward_modulated(x, weight, shift, scale, index) + assert torch.equal(ours, provider_norm_modulate(x, weight, shift, scale, index)) + + def test_norm_out_modulation_bitwise_equal_to_diffusers(self, h3_weights_cpu): + weight = h3_weights_cpu["norm_out.norm.weight"].cuda() + shift, scale = _table(3, 5376, 2, seed=3) # (T, H) rows indexed by timestep + ti, _ = h3_packed_layout(1000, 3, seed=3) + x = _x((1, 1000, 5376), seed=3) + ours = _cuda_op().forward_modulated(x, weight, shift, scale, ti) + assert torch.equal(ours, provider_norm_modulate(x, weight, shift, scale, ti)) + + @pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) + def test_modulation_bitwise_in_every_dtype(self, dtype): + """FP32 must not contract the separately rounded eager ops into an FMA.""" + + w = (torch.rand(5376, device="cuda") + 0.5).to(dtype) + shift, scale, *_ = _table(9, 5376, 6, dtype=dtype, seed=9) + ti, tags = h3_packed_layout(333, 3, seed=9) + index = ti * 3 + tags + x = _x((2, 333, 5376), dtype, seed=9) + ours = _cuda_op().forward_modulated(x, w, shift, scale, index) + assert torch.equal(ours, provider_norm_modulate(x, w, shift, scale, index)) + + def test_rows_are_batch_and_position_invariant(self): + w = (torch.rand(5376, device="cuda") + 0.5).bfloat16() + shift, scale, *_ = _table(9, 5376, 6, seed=4) + ti, tags = h3_packed_layout(300, 3, seed=4) + index = ti * 3 + tags + x = _x((3, 300, 5376), seed=4) + full = _cuda_op().forward_modulated(x, w, shift, scale, index) + part = _cuda_op().forward_modulated(x[2:3, 50:90], w, shift, scale, index[50:90]) + assert torch.equal(part[0], full[2, 50:90]) + + def test_rejects_unsupported(self): + w = torch.ones(6, device="cuda", dtype=torch.bfloat16) + with pytest.raises(RuntimeError, match="multiple of 4"): + _cuda_op()(_x((2, 6)), w) + with pytest.raises(ValueError): + _cuda_op()(_x((2, 8)).cpu(), torch.ones(8).bfloat16()) + + +_DEVICE_MESSAGES = { + "weight": "with x's dtype and device", + "index": r"row index must be a non-empty contiguous int64 tensor on cuda:0", + "rstd": "on x's device", +} + + +@requires_cuda +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="needs two CUDA devices") +class TestCudaNativeDeviceValidation: + @pytest.mark.parametrize("argument", ["weight", "index"]) + def test_forward_rejects_different_device(self, argument): + from rl_engine.backends.extension import _C + + _cuda_op() + x = torch.ones(2, 8, device="cuda:0", dtype=torch.bfloat16) + weight = torch.ones(8, device="cuda:0", dtype=x.dtype) + shift, scale = torch.zeros(2, 16, device="cuda:0", dtype=x.dtype).chunk(2, dim=-1) + index = torch.tensor([0, 1], device="cuda:0") + if argument == "weight": + weight = weight.to("cuda:1") + else: + index = index.to("cuda:1") + with pytest.raises(RuntimeError, match=_DEVICE_MESSAGES[argument]): + _C.h3_rmsnorm_forward(x, weight, 1e-5, shift, scale, index) + torch.cuda.synchronize(0) + torch.cuda.synchronize(1) + + @pytest.mark.parametrize("argument", ["weight", "index", "rstd"]) + def test_backward_rejects_different_device(self, argument): + from rl_engine.backends.extension import _C + + _cuda_op() + x = torch.ones(2, 8, device="cuda:0", dtype=torch.bfloat16) + weight = torch.ones(8, device="cuda:0", dtype=x.dtype) + shift, scale = torch.zeros(2, 16, device="cuda:0", dtype=x.dtype).chunk(2, dim=-1) + index = torch.tensor([0, 1], device="cuda:0") + _, rstd = _C.h3_rmsnorm_forward(x, weight, 1e-5, shift, scale, index) + if argument == "weight": + weight = weight.to("cuda:1") + elif argument == "index": + index = index.to("cuda:1") + else: + rstd = rstd.to("cuda:1") + with pytest.raises(RuntimeError, match=_DEVICE_MESSAGES[argument]): + _C.h3_rmsnorm_backward(torch.ones_like(x), x, weight, rstd, shift, scale, index) + torch.cuda.synchronize(0) + torch.cuda.synchronize(1) + + +@requires_cuda +class TestCudaBackward: + def _grads(self, fn, tensors, grad): + leaves = [t.detach().clone().requires_grad_(True) for t in tensors] + fn(*leaves).backward(grad) + return [leaf.grad for leaf in leaves] + + def test_modulated_backward_against_fp64(self): + g = torch.Generator(device="cpu").manual_seed(5) + w = (torch.rand(5376, generator=g) + 0.5).bfloat16().cuda() + table = (torch.randn(9, 2 * 5376, generator=g) * 0.5).bfloat16().cuda() + ti, tags = h3_packed_layout(2049, 3, seed=5) + index = ti * 3 + tags + x = _x((2, 2049, 5376), seed=5) + grad = _x((2, 2049, 5376), seed=6, scale=1.0) + op = _cuda_op() + + def ours(x_, w_, t_): + shift, scale = t_.chunk(2, dim=-1) + return op.forward_modulated(x_, w_, shift, scale, index) + + def golden(x_, w_, t_): + shift, scale = t_.chunk(2, dim=-1) + return _fp64_modulated(x_, w_, shift, scale, index) + + got = self._grads(ours, (x, w, table), grad) + ref = self._grads(golden, (x.double(), w.double(), table.double()), grad.double()) + for name, g_, r in zip(("dx", "dweight", "dtable"), got, ref): + rel = ((g_.double() - r).abs().max() / r.abs().max()).item() + # Two BF16 roundings remain: the eager VJP uses round(1 + scale), and + # every gradient is rounded once to BF16 (2^-9 relative each). + assert rel < 1e-2, (name, rel) + again = self._grads(ours, (x, w, table), grad) + assert all(torch.equal(a, b) for a, b in zip(got, again)) + + @pytest.mark.parametrize("seq", [257, 1024, 4097]) + def test_modulated_gradients_meet_contract_against_golden(self, seq): + # Against forward_modulated_fp32 before it stored norm(x) and 1 + scale + # in BF16, d_weight missed from S=1024 and d_scale at S=4097. + from rl_engine.contracts.numerical import load_contract, resolve_tolerance + + spec = resolve_tolerance( + load_contract(), + judgment="gradient_accuracy", + op_class="reduction", + dtype=torch.bfloat16, + ) + g = torch.Generator(device="cpu").manual_seed(seq) + w = (torch.rand(5376, generator=g) + 0.5).bfloat16().cuda() + table = (torch.randn(9, 2 * 5376, generator=g) * 0.1).bfloat16().cuda() + ti, tags = h3_packed_layout(seq, 3, seed=seq) + index = ti * 3 + tags + x = _x((1, seq, 5376), seed=seq) + grad = _x((1, seq, 5376), seed=seq + 1, scale=1.0) + op, native = _cuda_op(), NativeH3RMSNormOp() + + def run(fn, upstream): + def call(x_, w_, t_): + shift, scale = t_.chunk(2, dim=-1) + return fn(x_, w_, shift, scale, index) + + return self._grads(call, (x, w, table), upstream) + + got = run(op.forward_modulated, grad) + ref = run(native.forward_modulated_fp32, grad.float()) + for name, g_, r in zip(("dx", "dweight", "dtable"), got, ref): + torch.testing.assert_close( + g_.float(), r.float(), atol=spec.atol, rtol=spec.rtol, msg=lambda m: f"{name}: {m}" + ) + + def test_plain_backward_against_fp64_and_row_local_dx(self): + w = (torch.rand(5376, generator=torch.Generator().manual_seed(7)) + 0.5).bfloat16().cuda() + x = _x((4, 64, 5376), seed=7) + grad = _x((4, 64, 5376), seed=8, scale=1.0) + op = _cuda_op() + got = self._grads(lambda a, b: op(a, b), (x, w), grad) + + def golden(a, b): + return a * torch.rsqrt(a.square().mean(-1, keepdim=True) + 1e-5) * b + + ref = self._grads(golden, (x.double(), w.double()), grad.double()) + for g, r in zip(got, ref): + assert ((g.double() - r).abs().max() / r.abs().max()).item() < 5e-3 + part = self._grads(lambda a, b: op(a, b), (x[1:2, 10:20], w), grad[1:2, 10:20]) + assert torch.equal(part[0][0], got[0][1, 10:20]) + + def test_registry_dispatches_cuda(self): + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + assert ( + type(KernelRegistry().get_op("h3_rmsnorm", device="cuda")).__name__ == "H3RMSNormCudaOp" + ) diff --git a/tests/models/minimax_h3/test_h3_sp_norm_adaln.py b/tests/models/minimax_h3/test_h3_sp_norm_adaln.py new file mode 100644 index 000000000..4866e68f3 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_sp_norm_adaln.py @@ -0,0 +1,325 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``sp_norm_adaln``: sequence-parallel norm / modulation / gated residual. + +* ownership: contiguous position shards per batch item; bad layouts fail closed; +* backward (one GPU, ranks as threads around the plain backward functions): every + rank's row gradients equal WS1's rows, and its ``dweight``/``d_shift``/``d_scale``/ + ``d_gate`` equal WS1's, for SP 2..8, odd S, B > 1, interleaved modality rows and + bf16/fp32; +* end to end (real NCCL, autograd): the MLP-side block region + ``norm2(residual + gate_msa * y)`` with modulation is byte-equal to WS1 on every rank. +""" + +from __future__ import annotations + +import os +import threading +from types import SimpleNamespace + +import pytest +import torch + +from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + SPRowLayout, + sp_row_layout, +) +from rl_engine.validation.models.h3_ws2 import ( + run_world, + sp_case, + sp_matches_ws1, + sp_region_rank, + ws1_sp_region, +) + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +ALL_TRUE = dict.fromkeys(("out", "d_residual", "d_y", "d_weight", "d_table"), True) + + +def _require_ext(): + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + sp_norm_adaln_available, + ) + + if not sp_norm_adaln_available(): + pytest.skip("rl_engine._C lacks the H3 SP symbols") + + +class _ThreadGroup: + """Ranks as threads on one GPU; ``all_gather`` is the rank-ordered concatenation. + + Only for the plain backward functions: autograd runs every backward of a device + on one worker thread, so autograd-driven ranks need real processes. + """ + + def __init__(self, world: int) -> None: + self.barrier, self.slots = threading.Barrier(world), [None] * world + + def handle(self, rank: int): + def all_gather(t): + self.slots[rank] = t + self.barrier.wait() + out = torch.cat(self.slots) + self.barrier.wait() + return out + + return SimpleNamespace(rank=rank, world_size=len(self.slots), all_gather=all_gather) + + def run(self, fn): + results, errors = [None] * len(self.slots), [] + + def body(rank): + try: + results[rank] = fn(self.handle(rank)) + except BaseException as exc: # re-raised below + errors.append(exc) + self.barrier.abort() + + threads = [threading.Thread(target=body, args=(r,)) for r in range(len(self.slots))] + for t in threads: + t.start() + for t in threads: + t.join() + if errors: + raise next( + (e for e in errors if not isinstance(e, threading.BrokenBarrierError)), errors[0] + ) + return results + + +class TestLayout: + @pytest.mark.parametrize("seq, sp", [(8, 8), (257, 2), (1000, 3), (4097, 8)]) + def test_shards_cover_every_position_once(self, seq, sp): + bounds = [SPRowLayout(seq, 1, sp, r).bounds(r) for r in range(sp)] + assert bounds[0][0] == 0 and bounds[-1][1] == seq + assert all(a[1] == b[0] and a[0] < a[1] for a, b in zip(bounds, bounds[1:])) + + @pytest.mark.parametrize("seq, batch, sp, rank", [(3, 1, 4, 0), (8, 0, 2, 0), (8, 1, 2, 2)]) + def test_rejects(self, seq, batch, sp, rank): + with pytest.raises(ValueError): + sp_row_layout(seq, batch, sp, rank) + + +@requires_cuda +class TestBackwardRanks: + @pytest.mark.parametrize( + "sp, batch, seq, hidden, layout, dtype", + [ + (2, 1, 4097, 5376, "block", torch.bfloat16), + (2, 1, 4097, 5376, "interleaved", torch.bfloat16), + (3, 2, 1000, 256, "interleaved", torch.bfloat16), + (4, 2, 777, 512, "block", torch.float32), + (5, 1, 1301, 128, "interleaved", torch.bfloat16), + (8, 2, 2049, 256, "block", torch.bfloat16), + (8, 1, 40, 64, "interleaved", torch.float32), # shards far smaller than a tile + ], + ) + def test_ranks_reproduce_ws1(self, sp, batch, seq, hidden, layout, dtype): + _require_ext() + from rl_engine import _C + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + _segment_tiles, + ) + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + sp_gate_residual_backward, + sp_plan, + sp_rmsnorm_backward, + ) + + c = sp_case(batch, seq, hidden, layout=layout, dtype=dtype, seed=sp * 100 + seq) + chunks = c["table"].chunk(6, dim=-1) + shift, scale, gate = chunks[3], chunks[4], chunks[2] # shift/scale_mlp, gate_msa + x2, g2, y2 = (c[k].view(-1, hidden) for k in ("residual", "grad", "y")) + idx, rows = c["index"], shift.shape[0] + tiles = _segment_tiles(idx.repeat(batch), rows) + _, rstd = _C.h3_rmsnorm_forward(x2, c["weight"], 1e-5, shift, scale, idx) + ws1_norm = _C.h3_rmsnorm_backward(g2, x2, c["weight"], rstd, shift, scale, idx, *tiles) + ws1_gate = _C.h3_gate_residual_backward(g2, y2, gate, idx, *tiles) + + def rank(coll): + lay = sp_row_layout(seq, batch, sp, coll.rank) + part = slice(lay.lo, lay.hi) + local = { + k: c[k][:, part].reshape(-1, hidden).contiguous() for k in ("residual", "grad", "y") + } + rs = rstd.view(batch, seq)[:, part].reshape(-1).contiguous() + plan = sp_plan(lay, idx, rows) + norm = sp_rmsnorm_backward( + local["grad"], local["residual"], c["weight"], rs, shift, scale, plan, coll + ) + gate_grads = sp_gate_residual_backward(local["grad"], local["y"], gate, plan, coll) + return lay, norm, gate_grads + + for lay, (dx, dw, dsh, dsc), (dy, dgate) in _ThreadGroup(sp).run(rank): + part = slice(lay.lo, lay.hi) + assert torch.equal( + dx.view(batch, -1, hidden), ws1_norm[0].view(batch, seq, hidden)[:, part] + ) + assert torch.equal(dw, ws1_norm[1]) + assert torch.equal(dsh, ws1_norm[2].to(dtype)) and torch.equal( + dsc, ws1_norm[3].to(dtype) + ) + assert torch.equal( + dy.view(batch, -1, hidden), ws1_gate[0].view(batch, seq, hidden)[:, part] + ) + assert torch.equal(dgate, ws1_gate[1]) + + def test_plain_norm_ranks_reproduce_ws1(self): + _require_ext() + from rl_engine import _C + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + sp_plan, + sp_rmsnorm_backward, + ) + + c = sp_case(2, 1500, 256, seed=5) + x2, g2 = c["residual"].view(-1, 256), c["grad"].view(-1, 256) + _, rstd = _C.h3_rmsnorm_forward(x2, c["weight"], 1e-5) + ws1 = _C.h3_rmsnorm_backward(g2, x2, c["weight"], rstd) + + def rank(coll): + lay = sp_row_layout(1500, 2, 4, coll.rank) + part = slice(lay.lo, lay.hi) + loc = [c[k][:, part].reshape(-1, 256).contiguous() for k in ("grad", "residual")] + rs = rstd.view(2, 1500)[:, part].reshape(-1).contiguous() + return sp_rmsnorm_backward(*loc, c["weight"], rs, None, None, sp_plan(lay), coll) + + assert all(torch.equal(dw, ws1[1]) for _, dw, _, _ in _ThreadGroup(4).run(rank)) + + def test_fails_closed(self): + _require_ext() + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import ( + H3SPNormAdaLNCudaOp, + ) + + c = sp_case(1, 100, 64) + shift, scale = c["table"].chunk(6, dim=-1)[3:5] + op = H3SPNormAdaLNCudaOp(SimpleNamespace(rank=1, world_size=2), seq_len=100, batch=1) + local = c["residual"][:, 50:] + with pytest.raises(ValueError): # the full sequence instead of this rank's rows + op.norm_modulated(c["residual"], c["weight"], shift, scale, c["index"]) + with pytest.raises(ValueError): # a shard of the index instead of the full index + op.norm_modulated(local, c["weight"], shift, scale, c["index"][50:]) + with pytest.raises(IndexError): + op.gate_residual(local, c["y"][:, 50:], shift, c["index"] + 9) + with pytest.raises(ValueError): + op.norm(local.cpu(), c["weight"].cpu()) + assert op.readback()["positions"] == [50, 100] + + +@pytest.mark.parametrize("world", [2, 4, 8]) +@pytest.mark.skipif(int(os.environ.get("WORLD_SIZE", "1")) != 1, reason="owns its processes") +def test_nccl_region_byte_equal_to_ws1(world): + if torch.cuda.device_count() < world: + pytest.skip(f"needs {world} GPUs") + _require_ext() + case = sp_case(2, 4097, layout="interleaved", seed=world) + ws1 = {k: v.cpu() for k, v in ws1_sp_region(**case).items()} + torch.cuda.empty_cache() + ranks = run_world(world, sp_region_rank, {k: v.cpu() for k, v in case.items()}) + assert sp_matches_ws1(ws1, ranks) == ALL_TRUE + assert {r["readback"]["collective_backend"] for r in ranks} == {"cuda_ipc_fixed_tree"} + + +@requires_cuda +def test_plan_snapshots_caller_index(): + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import sp_plan + + index = torch.tensor([0, 1, 0, 1], device="cuda") + plan = sp_plan(sp_row_layout(4, 1, 1, 0), index, 2) + local, exchanged = plan.local_index.clone(), plan.index_buf.clone() + index.fill_(99) + assert torch.equal(plan.local_index, local) + assert torch.equal(plan.index_buf, exchanged) + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="needs two CUDA devices") +def test_plain_norm_uses_activation_device_and_separate_cache(): + _require_ext() + from rl_engine.backends.cuda.model_specific.minimax_h3.sp_norm_adaln import H3SPNormAdaLNCudaOp + + collective = SimpleNamespace(rank=0, world_size=1, all_gather=lambda x: x) + op = H3SPNormAdaLNCudaOp(collective, 5, 1) + with torch.cuda.device(0): + for device in ("cuda:1", "cuda:0"): + x = torch.randn(1, 5, 8, device=device, requires_grad=True) + weight = torch.randn(8, device=device, requires_grad=True) + ref_x = x.detach().clone().requires_grad_(True) + ref_w = weight.detach().clone().requires_grad_(True) + result = op.norm(x, weight) + expected = torch.nn.functional.rms_norm(ref_x, (8,), ref_w, 1e-5) + result.sum().backward() + expected.sum().backward() + torch.testing.assert_close(result, expected) + torch.testing.assert_close(x.grad, ref_x.grad) + torch.testing.assert_close(weight.grad, ref_w.grad) + assert len(op._plans) == 2 + assert op.plan(None, None, device="cuda:1").plan.send_local.device.index == 1 + assert op.plan(None, None, device="cuda:0").plan.send_local.device.index == 0 + + +def _worker_exit_without_result(rank, world, init, queue, path, target, kwargs): + import time + + if rank == 0: + os._exit(kwargs["exit_code"]) + time.sleep(60) + + +def _worker_reports_error(rank, world, init, queue, path, target, kwargs): + import time + + if rank == 0: + queue.put({"rank": rank, "error": "intentional worker failure"}) + time.sleep(60) + + +def _worker_never_reports(rank, world, init, queue, path, target, kwargs): + import time + + time.sleep(60) + + +@pytest.mark.parametrize("mode", ["clean_exit", "crash", "reported_error", "timeout"]) +def test_run_world_cleans_up_failed_workers(monkeypatch, mode): + import time + import torch.multiprocessing as mp + from rl_engine.validation.models import h3_ws2 + + workers = { + "clean_exit": _worker_exit_without_result, + "crash": _worker_exit_without_result, + "reported_error": _worker_reports_error, + "timeout": _worker_never_reports, + } + monkeypatch.setattr(h3_ws2, "_bootstrap", workers[mode]) + monkeypatch.setattr(torch.cuda, "device_count", lambda: 2) + before = {p.pid for p in mp.active_children()} + start = time.monotonic() + expected = TimeoutError if mode == "timeout" else RuntimeError + with pytest.raises(expected): + run_world( + 2, + None, + {}, + timeout=1 if mode == "timeout" else 30, + exit_code=7 if mode == "crash" else 0, + ) + assert time.monotonic() - start < 20 + assert {p.pid for p in mp.active_children()} <= before + + +def _worker_reports_result(rank, world, init, queue, path, target, kwargs): + torch.save({"value": torch.tensor(rank)}, path.with_name(f"rank{rank}.pt")) + queue.put({"rank": rank}) + + +def test_run_world_returns_rank_ordered_results(monkeypatch): + from rl_engine.validation.models import h3_ws2 + + monkeypatch.setattr(h3_ws2, "_bootstrap", _worker_reports_result) + monkeypatch.setattr(torch.cuda, "device_count", lambda: 2) + results = run_world(2, None, {}, timeout=30) + assert [result["rank"] for result in results] == [0, 1] + assert [result["value"].item() for result in results] == [0, 1] diff --git a/tests/models/minimax_h3/test_h3_timestep_mlp.py b/tests/models/minimax_h3/test_h3_timestep_mlp.py new file mode 100644 index 000000000..555a28ebe --- /dev/null +++ b/tests/models/minimax_h3/test_h3_timestep_mlp.py @@ -0,0 +1,296 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``timestep_mlp_fp32``: FP32 256 -> 5376 -> 2688 timestep MLP. + +* accuracy: forward and every gradient against an FP64 golden, on the pinned + checkpoint weights and on synthetic non-H3 sizes; +* invariance: a timestep's ``temb`` bytes (and its row-local ``dx``) do not + depend on how many timesteps share the call or where it sits; +* determinism: repeated forward/backward runs are bitwise equal; +* contract: BF16 inputs (RFC probe H7) and malformed shapes fail closed. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.reference.minimax_h3.timestep_mlp import NativeH3TimestepMLPOp +from rl_engine.reference.minimax_h3.timestep_sinusoid import NativeH3TimestepSinusoidOp +from rl_engine.validation.models.h3_cases import h3_timesteps + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +NAMES = ("x", "w1", "b1", "w2", "b2") + + +def _cuda_op(): + """Construct the CUDA MLP operator, skipping missing native linear kernels.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import det_linear_available + from rl_engine.backends.cuda.model_specific.minimax_h3.timestep_mlp import H3TimestepMLPCudaOp + + if not det_linear_available(): + pytest.skip("rl_engine._C lacks h3_det_linear_*") + return H3TimestepMLPCudaOp() + + +def _features(num, device="cuda", seed=0): + """Generate seeded H3 sinusoidal features on the requested test device.""" + + return NativeH3TimestepSinusoidOp().forward(h3_timesteps(num, seed=seed, device=device)) + + +def _synthetic_params(k_in, hidden, out, device="cuda", seed=0): + """Create seeded FP32 parameters for a two-layer SiLU MLP of the requested sizes.""" + + g = torch.Generator(device="cpu").manual_seed(seed) + w1 = torch.randn(hidden, k_in, generator=g) / k_in**0.5 + b1 = torch.randn(hidden, generator=g) * 0.1 + w2 = torch.randn(out, hidden, generator=g) / hidden**0.5 + b2 = torch.randn(out, generator=g) * 0.1 + return [t.to(device) for t in (w1, b1, w2, b2)] + + +def _real_params(weights): + """Return pinned timestep MLP weights and biases on CUDA in layer order.""" + + return [ + weights[f"time_embedder.linear_{i}.{p}"].cuda() for i in (1, 2) for p in ("weight", "bias") + ] + + +def _fp64_grads(x, params, grad_out): + """Compute gradients of both linear layers and SiLU using FP64 autograd.""" + + leaves = [t.detach().double().requires_grad_(True) for t in (x, *params)] + z = F.linear(leaves[0], leaves[1], leaves[2]) + F.linear(z * torch.sigmoid(z), leaves[3], leaves[4]).backward(grad_out.double()) + return [leaf.grad for leaf in leaves] + + +class TestReference: + def test_provider_and_golden_agree(self): + """Keep the FP32 provider close to the high-precision MLP replay on CPU.""" + + x = _features(5, device="cpu") + params = _synthetic_params(256, 64, 32, device="cpu") + op = NativeH3TimestepMLPOp() + torch.testing.assert_close( + op.forward(x, *params), op.forward_fp32(x, *params), atol=1e-5, rtol=1e-5 + ) + + @pytest.mark.parametrize("which", range(5)) + def test_rejects_bf16_anywhere(self, which): + # RFC probe H7: casting a declared FP32 path to BF16 early. + """Reject BF16 in every argument of the declared FP32 MLP path.""" + + args = [_features(2, device="cpu"), *_synthetic_params(256, 64, 32, device="cpu")] + args[which] = args[which].bfloat16() + with pytest.raises(TypeError, match="float32"): + NativeH3TimestepMLPOp().forward(*args) + + def test_rejects_shape_mismatches(self): + """Reject incompatible layer dimensions, malformed biases, and empty feature batches.""" + + x = _features(2, device="cpu") + w1, b1, w2, b2 = _synthetic_params(256, 64, 32, device="cpu") + op = NativeH3TimestepMLPOp() + with pytest.raises(ValueError): + op.forward(x[:, :128], w1, b1, w2, b2) + with pytest.raises(ValueError): + op.forward(x, w1, b1[:10], w2, b2) + with pytest.raises(ValueError): + op.forward(x, w1, b1, w2[:, :10], b2) + with pytest.raises(ValueError): + op.forward(x[:0], w1, b1, w2, b2) + + +@requires_cuda +class TestCudaRealWeights: + @pytest.mark.parametrize("num", [1, 2, 3, 4, 7]) + def test_forward_against_fp64_golden(self, h3_weights_cpu, num): + """Check checkpoint forward accuracy, FP32 dtype, and the timestep output shape.""" + + params = _real_params(h3_weights_cpu) + x = _features(num, seed=num) + out = _cuda_op()(x, *params) + gold = NativeH3TimestepMLPOp().forward_fp32(x, *params) + assert out.dtype == torch.float32 and out.shape == (num, 2688) + # tolerance_contract.json forward_accuracy / reduction / float32. + torch.testing.assert_close(out, gold, atol=1e-4, rtol=1e-4) + assert (out - gold).abs().max().item() < 1e-5 + + def test_rows_are_batch_and_position_invariant(self, h3_weights_cpu): + """Keep each checkpoint embedding row bitwise unchanged across batch sizes and positions.""" + + op = _cuda_op() + params = _real_params(h3_weights_cpu) + x_all = _features(9, seed=1) + full = op(x_all, *params) + for i in range(9): + assert torch.equal(op(x_all[i : i + 1], *params)[0], full[i]) + for num in (2, 3, 4, 5, 8): + for pos in (0, num - 1): + x = _features(num, seed=10 + num) + x[pos] = x_all[4] + assert torch.equal(op(x, *params)[pos], full[4]), (num, pos) + + def test_backward_against_fp64_golden(self, h3_weights_cpu): + """Match all five checkpoint input gradients to the FP64 autograd golden.""" + + params = _real_params(h3_weights_cpu) + x = _features(3, seed=2) + leaves = [t.detach().clone().requires_grad_(True) for t in (x, *params)] + grad_out = torch.randn( + 3, 2688, device="cuda", generator=torch.Generator("cuda").manual_seed(0) + ) + _cuda_op()(*leaves).backward(grad_out) + for name, leaf, ref in zip(NAMES, leaves, _fp64_grads(x, params, grad_out), strict=True): + # tolerance_contract.json gradient_accuracy / reduction / float32. + torch.testing.assert_close(leaf.grad.double(), ref, atol=1e-4, rtol=1e-4, msg=name) + + +@requires_cuda +class TestCudaSynthetic: + @pytest.mark.parametrize("k_in, hidden, out", [(8, 24, 16), (12, 40, 20), (256, 5376, 2688)]) + def test_other_sizes(self, k_in, hidden, out): + """Match the high-precision forward golden for small and checkpoint-sized MLPs.""" + + params = _synthetic_params(k_in, hidden, out, seed=k_in) + x = torch.randn(5, k_in, device="cuda") + torch.testing.assert_close( + _cuda_op()(x, *params), + NativeH3TimestepMLPOp().forward_fp32(x, *params), + atol=1e-4, + rtol=1e-4, + ) + + def test_zero_hidden_size_matches_cpu_forward_backward(self): + """Match CPU forward values and all input gradients with a zero-width hidden layer.""" + + generator = torch.Generator(device="cpu").manual_seed(0) + cpu_inputs = [ + torch.randn(3, 8, generator=generator), + torch.empty(0, 8), + torch.empty(0), + torch.empty(16, 0), + torch.linspace(-1.0, 1.0, 16), + ] + cpu_inputs = [tensor.requires_grad_(True) for tensor in cpu_inputs] + cuda_inputs = [tensor.detach().cuda().requires_grad_(True) for tensor in cpu_inputs] + cpu_out = NativeH3TimestepMLPOp()(*cpu_inputs) + cuda_out = _cuda_op()(*cuda_inputs) + torch.testing.assert_close(cuda_out.cpu(), cpu_out, atol=0, rtol=0) + grad = torch.randn(3, 16, generator=generator) + cpu_out.backward(grad) + cuda_out.backward(grad.cuda()) + for name, cpu_input, cuda_input in zip(NAMES, cpu_inputs, cuda_inputs, strict=True): + torch.testing.assert_close( + cuda_input.grad.cpu(), cpu_input.grad, atol=0, rtol=0, msg=name + ) + + def test_rejects_unaligned_k(self): + """Reject inner dimensions that violate the native FP32 four-element alignment.""" + + params = _synthetic_params(6, 8, 8) + with pytest.raises(RuntimeError, match="multiple of 4"): + _cuda_op()(torch.randn(2, 6, device="cuda"), *params) + + def test_rejects_cpu_and_mixed_devices(self): + """Reject CPU-only and mixed-device inputs at the CUDA MLP boundary.""" + + params = _synthetic_params(8, 24, 16) + with pytest.raises(ValueError): + _cuda_op()(torch.randn(2, 8), *[p.cpu() for p in params]) + with pytest.raises(ValueError): + _cuda_op()(torch.randn(2, 8), *params) + + def test_repeat_and_backward_are_deterministic(self): + """Require repeated MLP outputs and all input gradients to be bitwise equal.""" + + op = _cuda_op() + params = _synthetic_params(256, 5376, 2688, seed=3) + x = _features(4, seed=3) + grad_out = torch.randn(4, 2688, device="cuda") + runs = [] + for _ in range(3): + leaves = [t.detach().clone().requires_grad_(True) for t in (x, *params)] + out = op(*leaves) + out.backward(grad_out) + runs.append([out.detach(), *[leaf.grad for leaf in leaves]]) + for later in runs[1:]: + for a, b in zip(runs[0], later, strict=True): + assert torch.equal(a, b) + + def test_dx_rows_are_batch_invariant(self): + """Keep feature gradients bitwise unchanged when rows are evaluated independently.""" + + op = _cuda_op() + params = _synthetic_params(256, 5376, 2688, seed=4) + x = _features(6, seed=4) + grad_out = torch.randn(6, 2688, device="cuda") + + def dx(rows): + """Compute selected feature-row gradients using matching upstream rows.""" + + leaf = x[rows].detach().clone().requires_grad_(True) + op(leaf, *params).backward(grad_out[rows]) + return leaf.grad + + full = dx(slice(0, 6)) + for i in range(6): + assert torch.equal(dx(slice(i, i + 1))[0], full[i]) + + def test_registry_dispatches_cuda(self): + """Resolve the timestep MLP registry entry to its dedicated CUDA implementation.""" + + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("timestep_mlp_fp32", device="cuda") + assert type(op).__name__ == "H3TimestepMLPCudaOp" + + +@requires_cuda +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="needs two CUDA devices") +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +class TestCudaDetLinearDeviceContract: + def test_forward_rejects_bias_on_another_device(self, dtype): + """Reject a projection bias on a different CUDA device from its weight.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import det_linear_forward + + _cuda_op() + x = torch.randn(2, 8, device="cuda:0", dtype=dtype) + weight = torch.randn(16, 8, device="cuda:0", dtype=dtype) + bias = torch.randn(16, device="cuda:1", dtype=dtype) + with pytest.raises(RuntimeError, match="bias and weight must be on the same device"): + det_linear_forward(x, weight, bias) + + def test_backward_input_rejects_weight_on_another_device(self, dtype): + """Reject native input-gradient operands on different CUDA devices.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + det_linear_backward_input, + ) + + _cuda_op() + grad = torch.randn(2, 16, device="cuda:0") + weight = torch.randn(16, 8, device="cuda:1", dtype=dtype) + with pytest.raises(RuntimeError, match="grad and weight must be on the same device"): + det_linear_backward_input(grad, weight, dtype) + + def test_backward_weight_rejects_input_on_another_device(self, dtype): + """Reject native parameter-gradient operands on different CUDA devices.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import ( + det_linear_backward_weight, + ) + + _cuda_op() + grad = torch.randn(2, 16, device="cuda:0") + x = torch.randn(2, 8, device="cuda:1", dtype=dtype) + with pytest.raises(RuntimeError, match="grad and x must be on the same device"): + det_linear_backward_weight(grad, x, dtype) diff --git a/tests/models/minimax_h3/test_h3_timestep_sinusoid.py b/tests/models/minimax_h3/test_h3_timestep_sinusoid.py new file mode 100644 index 000000000..1d08d9c02 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_timestep_sinusoid.py @@ -0,0 +1,247 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``timestep_sinusoid_h3``: FP32 [cos | sin] timestep features. + +Three properties, each checked separately: + +* provider parity: the CUDA kernel is bitwise equal to the diffusers + ``get_timestep_embedding`` replay on the same device (an elementwise op + with the same FP32 operation sequence); +* accuracy: within the FP32 elementwise contract of an FP64 golden; +* invariance: a row's bytes do not depend on how many timesteps share the + launch, their order, or repeated runs. +""" + +from __future__ import annotations + +import math + +import pytest +import torch + +from rl_engine.reference.minimax_h3.timestep_sinusoid import NativeH3TimestepSinusoidOp +from rl_engine.validation.models.h3_cases import h3_timesteps +from rl_engine.validation.models.h3_provider import provider_time_proj + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") + + +def _cuda_op(): + """Construct the CUDA sinusoid operator, skipping a missing native forward entrypoint.""" + + from rl_engine.backends.cuda.model_specific.minimax_h3.timestep_sinusoid import ( + H3TimestepSinusoidCudaOp, + h3_sinusoid_cuda_available, + ) + + if not h3_sinusoid_cuda_available(): + pytest.skip("rl_engine._C lacks h3_timestep_sinusoid_forward") + return H3TimestepSinusoidCudaOp() + + +class TestReference: + def test_layout_is_cos_then_sin(self): + """Check FP32 cosine-then-sine layout and the unit-frequency boundary values.""" + + t = torch.tensor([0.0, 1.0]) + out = NativeH3TimestepSinusoidOp().forward(t) + assert out.shape == (2, 256) and out.dtype == torch.float32 + assert torch.equal(out[0, :128], torch.ones(128)) + assert torch.equal(out[0, 128:], torch.zeros(128)) + # k = 0 has frequency 1: cos(1), sin(1) at t = 1. + assert out[1, 0].item() == pytest.approx(math.cos(1.0), abs=1e-7) + assert out[1, 128].item() == pytest.approx(math.sin(1.0), abs=1e-7) + + def test_provider_path_matches_diffusers_cpu(self): + """Require CPU sinusoid values to match the diffusers provider bitwise.""" + + t = h3_timesteps(9, device="cpu") + ours = NativeH3TimestepSinusoidOp().forward(t) + assert torch.equal(ours, provider_time_proj(t, 256)) + + def test_golden_close_to_provider_path(self): + """Keep the CPU provider within FP32 error of the high-precision sinusoid golden.""" + + t = h3_timesteps(33, device="cpu") + op = NativeH3TimestepSinusoidOp() + torch.testing.assert_close(op.forward(t), op.forward_fp32(t), atol=1e-6, rtol=0) + + @pytest.mark.parametrize( + "bad, error", + [ + (torch.tensor([[0.5]]), ValueError), # not 1-D + (torch.empty(0), ValueError), # empty + (torch.tensor([1, 0]), TypeError), # integer + (torch.tensor([0.5, 1.5]), ValueError), # H10: t outside [0, 1] + (torch.tensor([500.0]), ValueError), # H10: t * 1000 convention + (torch.tensor([float("nan")]), ValueError), + ], + ) + def test_rejects_invalid_timesteps(self, bad, error): + """Reject empty, malformed, non-floating, non-finite, or out-of-range timesteps.""" + + with pytest.raises(error): + NativeH3TimestepSinusoidOp().forward(bad) + + def test_rejects_odd_channels(self): + """Reject odd feature counts that cannot split into equal cosine and sine halves.""" + + with pytest.raises(ValueError): + NativeH3TimestepSinusoidOp().forward(torch.tensor([0.5]), num_channels=255) + + +@requires_cuda +class TestCuda: + @pytest.mark.parametrize("num", [1, 2, 3, 7, 64, 1000, 4097]) + def test_bitwise_equal_to_diffusers_cuda_path(self, num): + """Require CUDA sinusoid values to match both provider replays for varied batch sizes.""" + + t = h3_timesteps(num, seed=num) + out = _cuda_op()(t) + assert out.dtype == torch.float32 and out.shape == (num, 256) + assert torch.equal(out, provider_time_proj(t, 256)) + assert torch.equal(out, NativeH3TimestepSinusoidOp().forward(t)) + + @pytest.mark.parametrize("num_channels", [2, 8, 96, 256, 320]) + def test_other_channel_counts_match_reference(self, num_channels): + """Preserve bitwise reference parity across supported even channel counts.""" + + t = h3_timesteps(5) + out = _cuda_op()(t, num_channels=num_channels) + ref = NativeH3TimestepSinusoidOp().forward(t, num_channels=num_channels) + assert torch.equal(out, ref) + + def test_fp32_contract_against_fp64_golden(self): + """Bound CUDA sinusoid error against the high-precision golden by FP32 precision.""" + + t = h3_timesteps(257) + out = _cuda_op()(t) + gold = NativeH3TimestepSinusoidOp().forward_fp32(t) + # tolerance_contract.json forward_accuracy / elementwise / float32. + torch.testing.assert_close(out, gold, atol=1e-5, rtol=1e-5) + assert (out - gold).abs().max().item() <= 2.0**-23 + + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) + def test_low_precision_timesteps_are_upcast_first(self, dtype): + """Match provider FP32 upcasting for BF16 and FP16 timesteps.""" + + t = h3_timesteps(6).to(dtype) + assert torch.equal(_cuda_op()(t), provider_time_proj(t, 256)) + + def test_batch_position_and_repeat_invariance(self): + """Keep row bits unchanged across batch positions, repeated calls, and permutations.""" + + op = _cuda_op() + base = torch.tensor(0.3141592, device="cuda") + single = op(base.view(1)) + for num in (2, 5, 64, 1023): + for pos in (0, num // 2, num - 1): + t = h3_timesteps(num, seed=num) + t[pos] = base + assert torch.equal(op(t)[pos], single[0]), (num, pos) + t = h3_timesteps(129, seed=11) + first = op(t) + for _ in range(3): + assert torch.equal(op(t), first) + perm = torch.randperm(129, generator=torch.Generator().manual_seed(0)).cuda() + assert torch.equal(op(t[perm]), first[perm]) + + def test_non_contiguous_input(self): + """Treat strided timesteps identically to their contiguous copies.""" + + t = h3_timesteps(20) + strided = t[::2] + assert not strided.is_contiguous() + assert torch.equal(_cuda_op()(strided), _cuda_op()(strided.contiguous())) + + def test_rejects_cpu_input(self): + """Reject CPU timesteps at the CUDA sinusoid boundary.""" + + with pytest.raises(ValueError): + _cuda_op()(torch.tensor([0.5])) + + def test_rejects_out_of_range(self): + """Reject scaled timesteps outside the declared unit interval.""" + + with pytest.raises(ValueError): + _cuda_op()(torch.tensor([0.0, 999.0], device="cuda")) + + @pytest.mark.parametrize("bad", [-0.1, 1.1, float("nan"), float("inf"), -float("inf")]) + def test_native_entrypoint_rejects_invalid_values(self, bad): + """Reject invalid values in native and wrapper calls even when neighbors are valid.""" + + from rl_engine.backends.extension import _C + + _cuda_op() + # Valid neighbors must not hide an invalid timestep in the native call. + t = torch.tensor([0.0, bad, 1.0], device="cuda") + with pytest.raises(ValueError, match=r"finite and lie in \[0, 1\]"): + _C.h3_timestep_sinusoid_forward(t) + with pytest.raises(ValueError, match=r"finite and lie in \[0, 1\]"): + _cuda_op()(t) + + def test_native_entrypoint_accepts_boundaries(self): + """Accept both interval endpoints and match provider values in the native call.""" + + from rl_engine.backends.extension import _C + + _cuda_op() + t = torch.tensor([0.0, 0.5, 1.0], device="cuda") + assert torch.equal(_C.h3_timestep_sinusoid_forward(t), provider_time_proj(t, 256)) + + def test_trusted_path_supports_cuda_graph(self): + """Capture and replay the prevalidated sinusoid path without changing its output.""" + + op = _cuda_op() + t = h3_timesteps(7) + expected = op(t) + warmup = torch.cuda.Stream() + warmup.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(warmup): + op.forward(t, check_range=False) + torch.cuda.current_stream().wait_stream(warmup) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = op.forward(t, check_range=False) + graph.replay() + assert torch.equal(out, expected) + + def test_backward_matches_autograd_reference(self): + """Match timestep gradients to the high-precision autograd sinusoid replay.""" + + t = h3_timesteps(17).requires_grad_(True) + t_ref = t.detach().clone().requires_grad_(True) + grad = torch.randn(17, 256, device="cuda", generator=torch.Generator("cuda").manual_seed(1)) + _cuda_op()(t).backward(grad) + NativeH3TimestepSinusoidOp().forward_fp32(t_ref).backward(grad) + torch.testing.assert_close(t.grad, t_ref.grad, atol=1e-4, rtol=1e-4) + + def test_backward_is_row_invariant(self): + """Keep a timestep gradient bitwise unchanged when unrelated rows join the batch.""" + + op = _cuda_op() + grad_row = torch.randn(256, device="cuda", generator=torch.Generator("cuda").manual_seed(2)) + + def grad_of(t, pos): + """Differentiate one selected timestep using a fixed upstream feature-gradient row.""" + + t = t.clone().requires_grad_(True) + grad = torch.zeros(t.shape[0], 256, device="cuda") + grad[pos] = grad_row + op(t).backward(grad) + return t.grad[pos] + + single = grad_of(torch.tensor([0.7], device="cuda"), 0) + t = h3_timesteps(40) + t[13] = 0.7 + assert torch.equal(grad_of(t, 13), single) + + def test_registry_dispatches_cuda(self): + """Resolve the sinusoid registry entry to its dedicated CUDA implementation.""" + + from rl_engine.runtime.registry import KernelRegistry + + _cuda_op() + op = KernelRegistry().get_op("timestep_sinusoid_h3", device="cuda") + assert type(op).__name__ == "H3TimestepSinusoidCudaOp" diff --git a/tests/models/minimax_h3/test_h3_tp_adaln.py b/tests/models/minimax_h3/test_h3_tp_adaln.py new file mode 100644 index 000000000..c9b511153 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_tp_adaln.py @@ -0,0 +1,159 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""RFC #420 ``tp_adaln_3mod``: column-parallel AdaLN projection, byte-equal to WS1. + +* ownership: contiguous 64-aligned column shards cover the 3 x 6 x H table + exactly once; unsupported TP sizes and wrong shards fail closed; +* shard arithmetic (one GPU): a rank's forward columns and ``d_input`` chunk + partials are bitwise the matching slice of the full WS1 call, and folding + the rank-ordered partials reproduces WS1's ``d_input``; +* real NCCL (2/4/8 GPUs): every rank's table and ``d_temb`` and the + concatenated ``dW``/``db`` shards are bitwise equal to WS1. +""" + +from __future__ import annotations + +import os +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + +from rl_engine.backends.cuda.model_specific.minimax_h3.tp_adaln_projection import adaln_column_shard +from rl_engine.validation.models.h3_ws2 import ( + run_world, + tp_matches_ws1, + tp_projection_rank, + ws1_projection, +) + +requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs a CUDA device") +WEIGHT = "transformer_blocks.0.adaln_proj.linear.weight" +BIAS = "transformer_blocks.0.adaln_proj.linear.bias" +N_H3 = 3 * 6 * 5376 +ALL_TRUE = dict.fromkeys(("table", "d_temb", "d_weight", "d_bias"), True) + + +def _require_ext(): + from rl_engine.backends.cuda.model_specific.minimax_h3.det_linear import det_linear_available + + if not det_linear_available(): + pytest.skip("rl_engine._C lacks h3_det_linear_*") + + +def _inputs(num_t, n_out, k, dtype=torch.bfloat16, seed=0): + g = torch.Generator(device="cpu").manual_seed(seed) + temb = (torch.randn(num_t, k, generator=g) * 2).cuda() + weight = (torch.randn(n_out, k, generator=g) / k**0.5).to(dtype).cuda() + bias = (torch.randn(n_out, generator=g) * 0.1).to(dtype).cuda() + grad = torch.randn(num_t, n_out, generator=g).to(dtype).cuda() + return temb, weight, bias, grad + + +class TestOwnership: + @pytest.mark.parametrize("tp", [1, 2, 3, 4, 6, 7, 8, 9]) + def test_shards_tile_the_table_once(self, tp): + shards = [adaln_column_shard(N_H3, tp, r) for r in range(tp)] + assert [s.begin for s in shards] == [0, *[s.end for s in shards[:-1]]] + assert shards[-1].end == N_H3 + assert sum(h1 - h0 for s in shards for (_, _, h0, h1) in s.slots()) == N_H3 + + def test_tp2_slots(self): + # rank 0: video's six chunks and text's first three; rank 1: the rest. + assert adaln_column_shard(N_H3, 2, 0).slots() == [ + *[(0, c, 0, 5376) for c in range(6)], + *[(1, c, 0, 5376) for c in range(3)], + ] + + @pytest.mark.parametrize("tp, rank", [(5, 0), (16, 0), (0, 0), (2, 2), (2, -1)]) + def test_rejects_unsupported(self, tp, rank): + with pytest.raises(ValueError): + adaln_column_shard(N_H3, tp, rank) + + +@requires_cuda +class TestShardArithmetic: + @pytest.mark.parametrize("tp", [2, 4, 8]) + def test_pinned_shards_reproduce_ws1(self, h3_weights_cpu, tp): + _require_ext() + from rl_engine.backends.cuda.model_specific.minimax_h3 import det_linear as dl + + weight, bias = h3_weights_cpu[WEIGHT].cuda(), h3_weights_cpu[BIAS].cuda() + g = torch.Generator(device="cpu").manual_seed(tp) + act = F.silu(torch.randn(3, weight.shape[1], generator=g).cuda() * 2).bfloat16() + grad = torch.randn(3, weight.shape[0], generator=g).cuda() + (full,) = dl.det_linear_forward(act, weight, bias) + full_partials = dl.det_linear_backward_input_partials(grad, weight) + pieces = [] + for rank in range(tp): + s = adaln_column_shard(weight.shape[0], tp, rank) + (cols,) = dl.det_linear_forward(act, weight[s.begin : s.end], bias[s.begin : s.end]) + assert torch.equal(cols, full[:, s.begin : s.end]) + part = dl.det_linear_backward_input_partials( + grad[:, s.begin : s.end].contiguous(), weight[s.begin : s.end] + ) + c0 = s.begin // dl.DINPUT_CHUNK + assert torch.equal(part, full_partials[c0 : c0 + part.shape[0]]) + pieces.append(part) + folded = dl.det_linear_fold_chunks(torch.cat(pieces), torch.float32) + assert torch.equal(folded, dl.det_linear_backward_input(grad, weight, torch.float32)) + + def test_tp1_op_is_ws1(self): + _require_ext() + one = SimpleNamespace(rank=0, world_size=1, backend_id="none", all_gather=lambda t: t) + temb, weight, bias, grad = _inputs(3, 3 * 6 * 256, 136) + ws1 = ws1_projection(temb, weight, bias, grad) + assert tp_matches_ws1(ws1, [tp_projection_rank(one, temb, weight, bias, grad)]) == ALL_TRUE + + def test_fails_closed(self): + _require_ext() + from rl_engine.backends.cuda.model_specific.minimax_h3.tp_adaln_projection import ( + H3TPAdaLNProjectionCudaOp, + shard_adaln_projection, + ) + + temb, weight, bias, _ = _inputs(2, 3 * 6 * 256, 136) + rank1 = SimpleNamespace(rank=1, world_size=2, backend_id="unused", all_gather=None) + op = H3TPAdaLNProjectionCudaOp(rank1, weight.shape[0]) + w, b = shard_adaln_projection(weight, bias, 2, 1) + with pytest.raises(TypeError): # RFC probe H7: cast before the SiLU + op(temb.bfloat16(), w, b) + with pytest.raises(ValueError): # the full weight instead of this rank's rows + op(temb, weight, bias) + with pytest.raises(ValueError): + op(temb.cpu(), w, b) + with pytest.raises(ValueError): # 4608 / 16 is not a multiple of 64 + H3TPAdaLNProjectionCudaOp(SimpleNamespace(rank=0, world_size=16), weight.shape[0]) + assert op.readback()["columns"] == [2304, 4608] + + +@pytest.mark.parametrize("world", [2, 4, 8]) +@pytest.mark.skipif(int(os.environ.get("WORLD_SIZE", "1")) != 1, reason="owns its processes") +def test_nccl_byte_equal_to_ws1(h3_weights_cpu, world): + if torch.cuda.device_count() < world: + pytest.skip(f"needs {world} GPUs") + _require_ext() + weight, bias = h3_weights_cpu[WEIGHT], h3_weights_cpu[BIAS] + g = torch.Generator(device="cpu").manual_seed(world) + inputs = { + "temb": torch.randn(3, weight.shape[1], generator=g) * 2, + "weight": weight, + "bias": bias, + "grad": torch.randn(3, weight.shape[0], generator=g).bfloat16(), + } + ws1 = { + k: v.cpu() for k, v in ws1_projection(**{k: v.cuda() for k, v in inputs.items()}).items() + } + torch.cuda.empty_cache() + ranks = run_world(world, tp_projection_rank, inputs) + assert tp_matches_ws1(ws1, ranks) == ALL_TRUE + assert [r["readback"]["rank"] for r in ranks] == list(range(world)) + assert {r["readback"]["collective_backend"] for r in ranks} == {"cuda_ipc_fixed_tree"} + + +@pytest.mark.parametrize("n_total", [0, -1152]) +def test_rejects_nonpositive_projection(n_total): + with pytest.raises(ValueError, match="N must be positive"): + adaln_column_shard(n_total, 1, 0) diff --git a/tests/models/minimax_h3/test_h3_weights.py b/tests/models/minimax_h3/test_h3_weights.py new file mode 100644 index 000000000..13b9b2489 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_weights.py @@ -0,0 +1,74 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Extracted H3 weights must retain the pinned tensor contents.""" + +from __future__ import annotations + +import hashlib + +import pytest +import torch + +from rl_engine.validation.models import h3_weights + + +@pytest.mark.parametrize( + "dtype,raw_bytes", + [ + (torch.float32, bytes.fromhex("000000000000803f")), + (torch.bfloat16, bytes.fromhex("0000803f")), + ], +) +def test_tensor_checksum_hashes_raw_bytes(dtype, raw_bytes): + """Hash the exact FP32 and BF16 bit patterns without numerical conversion.""" + + tensor = torch.tensor([0.0, 1.0], dtype=dtype) + assert h3_weights.sha256_tensor(tensor) == hashlib.sha256(raw_bytes).hexdigest() + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("mismatch", [None, "dtype", "shape", "sha256"]) +@pytest.mark.parametrize("names", [None, ["bias"]]) +def test_loading_rejects_modified_weight_contents(tmp_path, monkeypatch, dtype, mismatch, names): + """Reject altered shape, dtype, or bytes even outside the requested tensor subset.""" + + from safetensors.torch import save_file + + tensors = { + "weight": torch.tensor([[0.0, 1.0], [2.0, 3.0]], dtype=dtype), + "bias": torch.tensor([0.0, 1.0], dtype=dtype), + } + manifest = { + "tensors": { + name: { + "dtype": str(tensor.dtype).removeprefix("torch."), + "shape": list(tensor.shape), + "sha256": h3_weights.sha256_tensor(tensor), + } + for name, tensor in tensors.items() + } + } + if mismatch == "dtype": + tensors["weight"] = tensors["weight"].to(torch.float64) + elif mismatch == "shape": + tensors["weight"] = tensors["weight"].reshape(4) + elif mismatch == "sha256": + tensors["weight"][0, 0] = 1.0 + # Tensor checksums do not depend on serializer metadata or tensor key order. + save_file( + dict(reversed(list(tensors.items()))), + str(tmp_path / h3_weights.EXTRACTED_FILE), + metadata={"revision": "a different serializer metadata value"}, + ) + monkeypatch.setenv(h3_weights.WEIGHTS_ENV, str(tmp_path)) + monkeypatch.setattr(h3_weights, "load_h3_manifest", lambda: manifest) + + if mismatch: + with pytest.raises(ValueError, match=f"weight: {mismatch} .* != manifest"): + h3_weights.load_h3_conditioning_weights("cpu", names=names) + else: + actual = h3_weights.load_h3_conditioning_weights("cpu", names=names) + assert list(actual) == (names if names is not None else list(tensors)) + for name, tensor in actual.items(): + assert torch.equal(tensor, tensors[name]) diff --git a/tests/models/minimax_h3/test_h3_ws2_processes.py b/tests/models/minimax_h3/test_h3_ws2_processes.py new file mode 100644 index 000000000..f448ba580 --- /dev/null +++ b/tests/models/minimax_h3/test_h3_ws2_processes.py @@ -0,0 +1,47 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""CPU process lifecycle regressions for the multi-GPU evidence runner.""" + +import os +import time + +import pytest +import torch +import torch.multiprocessing as mp + +from rl_engine.validation.models import h3_ws2 + + +def _worker(rank, world, init, queue, path, target, kwargs): + mode = kwargs["mode"] + if mode == "success": + torch.save({"value": torch.tensor(rank)}, path.with_name(f"rank{rank}.pt")) + queue.put({"rank": rank}) + elif rank == 0 and mode == "error": + queue.put({"rank": rank, "error": "original rank traceback"}) + elif rank == 0 and mode == "crash": + os._exit(7) + else: + time.sleep(120) + + +@pytest.mark.parametrize("mode", ["success", "error", "crash", "timeout"]) +def test_run_world_reaps_workers(monkeypatch, mode): + monkeypatch.setattr(torch.cuda, "device_count", lambda: 2) + monkeypatch.setattr(h3_ws2, "_bootstrap", _worker) + before = {p.pid for p in mp.active_children()} + start = time.monotonic() + if mode == "success": + ranks = h3_ws2.run_world(2, None, {}, timeout=30, mode=mode) + assert [r["rank"] for r in ranks] == [0, 1] + assert [r["value"].item() for r in ranks] == [0, 1] + elif mode == "timeout": + with pytest.raises(TimeoutError, match="timed out"): + h3_ws2.run_world(2, None, {}, timeout=0.2, mode=mode) + else: + message = "original rank traceback" if mode == "error" else "rank 0 exited with code 7" + with pytest.raises(RuntimeError, match=message): + h3_ws2.run_world(2, None, {}, timeout=30, mode=mode) + assert time.monotonic() - start < 25 + assert {p.pid for p in mp.active_children()} == before diff --git a/tests/models/minimax_h3/test_prepare_h3_weights.py b/tests/models/minimax_h3/test_prepare_h3_weights.py new file mode 100644 index 000000000..5c9a966b9 --- /dev/null +++ b/tests/models/minimax_h3/test_prepare_h3_weights.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Weight extraction rejects off-manifest tensors even under Python optimization.""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import pytest +import torch + +from rl_engine.validation.models.h3_weights import EXTRACTED_FILE, sha256_file, sha256_tensor + + +@pytest.mark.parametrize("optimize", [0, 1]) +@pytest.mark.parametrize("mismatch", [None, "dtype", "shape", "sha256"]) +def test_manifest_validation_before_write(tmp_path, monkeypatch, optimize, mismatch): + """Enforce tensor contracts before extraction writes, including under Python optimization.""" + + import huggingface_hub + from safetensors.torch import load_file, save_file + + tensor = torch.tensor([0.0, 1.0]) + config = tmp_path / "config.json" + config.write_text("{}") + index = tmp_path / "index.json" + index.write_text(json.dumps({"weight_map": {"weight": "shard.safetensors"}})) + shard = tmp_path / "shard.safetensors" + save_file({"weight": tensor}, str(shard)) + manifest = { + "model_identity": { + "hf_repo": "test/h3", + "subfolder": "dit", + "revision": "pinned", + "config_sha256": sha256_file(config), + "index_file": index.name, + "index_sha256": sha256_file(index), + }, + "weight_shards": {shard.name: {"sha256": sha256_file(shard)}}, + "tensors": { + "weight": { + "dtype": "bfloat16" if mismatch == "dtype" else "float32", + "shape": [3] if mismatch == "shape" else [2], + "sha256": "0" * 64 if mismatch == "sha256" else sha256_tensor(tensor), + } + }, + } + monkeypatch.setattr( + huggingface_hub, + "hf_hub_download", + lambda repo, filename, **kwargs: str(tmp_path / Path(filename).name), + ) + out_dir = tmp_path / "extracted" + monkeypatch.setattr(sys, "argv", ["prepare_h3_weights.py", "--out", str(out_dir)]) + script = Path(__file__).resolve().parents[3] / "tools" / "weights" / "prepare_h3_weights.py" + namespace = {"__name__": "prepare_h3_weights_test", "__file__": str(script)} + # Compiling the actual script at optimize=1 reproduces python -O's assert removal. + exec(compile(script.read_text(), str(script), "exec", optimize=optimize), namespace) + namespace["load_h3_manifest"] = lambda: manifest + target = out_dir / EXTRACTED_FILE + if mismatch: + with pytest.raises(SystemExit, match=f"weight: {mismatch} .* != manifest"): + namespace["main"]() + assert not target.exists() + else: + namespace["main"]() + assert torch.equal(load_file(str(target))["weight"], tensor) diff --git a/tests/runtime/test_dispatch.py b/tests/runtime/test_dispatch.py index 3c47bbbcf..84e2463e2 100644 --- a/tests/runtime/test_dispatch.py +++ b/tests/runtime/test_dispatch.py @@ -4,6 +4,7 @@ import sys from types import ModuleType +import pytest import torch import rl_engine.platforms.device as device_module @@ -231,6 +232,30 @@ def fake_load_backend(backend): ] +H3_OPS = ( + "timestep_sinusoid_h3", + "timestep_mlp_fp32", + "adaln_projection_3mod", + "adaln_row_gather", + "h3_rmsnorm", + "adaln_gate_residual", + "final_adaln_out", +) + + +@pytest.mark.parametrize("op_name", H3_OPS) +def test_h3_ops_prefer_cuda_and_fall_back_to_pytorch_reference(op_name): + """Prefer H3 CUDA backends on CUDA and use their PyTorch fallbacks elsewhere.""" + + registry = KernelRegistry() + cuda_chain = registry._priority_map["cuda"][op_name] + assert cuda_chain[0].name.startswith("CUDA_H3_") + assert cuda_chain[-1].name.startswith("PYTORCH_H3_") + for platform in ("rocm", "musa", "cpu", "npu"): + chain = registry._priority_map[platform][op_name] + assert [backend.name for backend in chain] == [cuda_chain[-1].name] + + def test_npu_available_handles_runtime_failure(monkeypatch): class BrokenNPU: @staticmethod diff --git a/tests/validation/operators/test_ws1_gtest_gpu.py b/tests/validation/operators/test_ws1_gtest_gpu.py index bcaebc2f5..e7d1e5275 100644 --- a/tests/validation/operators/test_ws1_gtest_gpu.py +++ b/tests/validation/operators/test_ws1_gtest_gpu.py @@ -33,6 +33,8 @@ def _run(script: str, *args: str, timeout: int = 300) -> None: def test_all_ws1_single_ops_are_registered(): + """Require gtest specs for all WS1 single operators and the H3 conditioning stages.""" + names = set(operator_names()) assert { "rms_norm", @@ -49,6 +51,13 @@ def test_all_ws1_single_ops_are_registered(): "pack", "linear_logp", "prefix_shared_attention", + "timestep_sinusoid_h3", + "timestep_mlp_fp32", + "adaln_projection_3mod", + "adaln_row_gather", + "h3_rmsnorm", + "adaln_gate_residual", + "final_adaln_out", } <= names diff --git a/tools/migration/README.md b/tools/migration/README.md index 4a6cee300..48491e383 100644 --- a/tools/migration/README.md +++ b/tools/migration/README.md @@ -8,3 +8,6 @@ Historical workload/operator identifiers and archived evidence intentionally retain their original values. New source imports should use the mapped canonical module. Remove compatibility paths only in a separately announced migration once downstream launchers, patches, imports and saved objects have migrated. + +The MiniMax-H3 entries also record the open RFC #420 PR stack migration to +`test-h3`; those H3 files were not part of the original main baseline. diff --git a/tools/migration/layout.json b/tools/migration/layout.json index cb8fe469a..8d469480c 100644 --- a/tools/migration/layout.json +++ b/tools/migration/layout.json @@ -347,6 +347,7 @@ "benchmarks/benchmark_deterministic_attention.py": "benchmarks/operators/attention/benchmark_deterministic_attention.py", "benchmarks/benchmark_grpo_loss.py": "benchmarks/operators/loss/benchmark_grpo_loss.py", "benchmarks/benchmark_grpo_op.py": "benchmarks/operators/loss/benchmark_grpo_op.py", + "benchmarks/benchmark_h3_conditioning.py": "benchmarks/models/benchmark_h3_conditioning.py", "benchmarks/benchmark_linear_logp.py": "benchmarks/operators/logprob/benchmark_linear_logp.py", "benchmarks/benchmark_pack.py": "benchmarks/operators/packing/benchmark_pack.py", "benchmarks/benchmark_qwen_ffn_layout.py": "benchmarks/layers/benchmark_qwen_ffn_layout.py", @@ -466,6 +467,40 @@ "docs/design/ws2-attention-single-gpu-harness.md": "docs/archive/design/ws2-attention-single-gpu-harness.md", "docs/design/ws2-cp-attention-contract.md": "docs/contracts/ws2-cp-attention-contract.md", "docs/design/ws2-module-debug-matrix.md": "docs/archive/design/ws2-module-debug-matrix.md", + "docs/usage/evidence/h3-adaln-gate-residual-b200/figure.png": "reports/experiments/h3-adaln-gate-residual-b200/figure.png", + "docs/usage/evidence/h3-adaln-gate-residual-b200/report.json": "reports/experiments/h3-adaln-gate-residual-b200/report.json", + "docs/usage/evidence/h3-adaln-projection-b200/figure.png": "reports/experiments/h3-adaln-projection-b200/figure.png", + "docs/usage/evidence/h3-adaln-projection-b200/report.json": "reports/experiments/h3-adaln-projection-b200/report.json", + "docs/usage/evidence/h3-adaln-row-gather-b200/chain_replay.json": "reports/experiments/h3-adaln-row-gather-b200/chain_replay.json", + "docs/usage/evidence/h3-adaln-row-gather-b200/figure.png": "reports/experiments/h3-adaln-row-gather-b200/figure.png", + "docs/usage/evidence/h3-adaln-row-gather-b200/report.json": "reports/experiments/h3-adaln-row-gather-b200/report.json", + "docs/usage/evidence/h3-final-adaln-out-b200/chain_replay.json": "reports/experiments/h3-final-adaln-out-b200/chain_replay.json", + "docs/usage/evidence/h3-final-adaln-out-b200/figure.png": "reports/experiments/h3-final-adaln-out-b200/figure.png", + "docs/usage/evidence/h3-final-adaln-out-b200/report.json": "reports/experiments/h3-final-adaln-out-b200/report.json", + "docs/usage/evidence/h3-prior-art-b200/adaln_projection.json": "reports/experiments/h3-prior-art-b200/adaln_projection.json", + "docs/usage/evidence/h3-prior-art-b200/adaln_projection.png": "reports/experiments/h3-prior-art-b200/adaln_projection.png", + "docs/usage/evidence/h3-prior-art-b200/adaln_row_gather.json": "reports/experiments/h3-prior-art-b200/adaln_row_gather.json", + "docs/usage/evidence/h3-prior-art-b200/adaln_row_gather.png": "reports/experiments/h3-prior-art-b200/adaln_row_gather.png", + "docs/usage/evidence/h3-prior-art-b200/final_adaln_out.json": "reports/experiments/h3-prior-art-b200/final_adaln_out.json", + "docs/usage/evidence/h3-prior-art-b200/final_adaln_out.png": "reports/experiments/h3-prior-art-b200/final_adaln_out.png", + "docs/usage/evidence/h3-prior-art-b200/gate_residual.json": "reports/experiments/h3-prior-art-b200/gate_residual.json", + "docs/usage/evidence/h3-prior-art-b200/gate_residual.png": "reports/experiments/h3-prior-art-b200/gate_residual.png", + "docs/usage/evidence/h3-prior-art-b200/norm_modulate.json": "reports/experiments/h3-prior-art-b200/norm_modulate.json", + "docs/usage/evidence/h3-prior-art-b200/norm_modulate.png": "reports/experiments/h3-prior-art-b200/norm_modulate.png", + "docs/usage/evidence/h3-prior-art-b200/timestep_mlp.json": "reports/experiments/h3-prior-art-b200/timestep_mlp.json", + "docs/usage/evidence/h3-prior-art-b200/timestep_mlp.png": "reports/experiments/h3-prior-art-b200/timestep_mlp.png", + "docs/usage/evidence/h3-prior-art-b200/timestep_sinusoid.json": "reports/experiments/h3-prior-art-b200/timestep_sinusoid.json", + "docs/usage/evidence/h3-prior-art-b200/timestep_sinusoid.png": "reports/experiments/h3-prior-art-b200/timestep_sinusoid.png", + "docs/usage/evidence/h3-rmsnorm-b200/figure.png": "reports/experiments/h3-rmsnorm-b200/figure.png", + "docs/usage/evidence/h3-rmsnorm-b200/report.json": "reports/experiments/h3-rmsnorm-b200/report.json", + "docs/usage/evidence/h3-sp-norm-adaln-b200/figure.png": "reports/experiments/h3-sp-norm-adaln-b200/figure.png", + "docs/usage/evidence/h3-sp-norm-adaln-b200/report.json": "reports/experiments/h3-sp-norm-adaln-b200/report.json", + "docs/usage/evidence/h3-timestep-mlp-b200/figure.png": "reports/experiments/h3-timestep-mlp-b200/figure.png", + "docs/usage/evidence/h3-timestep-mlp-b200/report.json": "reports/experiments/h3-timestep-mlp-b200/report.json", + "docs/usage/evidence/h3-timestep-sinusoid-b200/figure.png": "reports/experiments/h3-timestep-sinusoid-b200/figure.png", + "docs/usage/evidence/h3-timestep-sinusoid-b200/report.json": "reports/experiments/h3-timestep-sinusoid-b200/report.json", + "docs/usage/evidence/h3-tp-adaln-3mod-b200/figure.png": "reports/experiments/h3-tp-adaln-3mod-b200/figure.png", + "docs/usage/evidence/h3-tp-adaln-3mod-b200/report.json": "reports/experiments/h3-tp-adaln-3mod-b200/report.json", "envs.py": "build_tools/envs.py", "examples/cross_config_qwen3_8b_megatron_tp2_cp2_vllm.json": "configs/experiments/cross_config/cross_config_qwen3_8b_megatron_tp2_cp2_vllm.json", "examples/cross_config_s0_cpu_smoke.json": "configs/experiments/cross_config/cross_config_s0_cpu_smoke.json", @@ -675,6 +710,19 @@ "rl_engine/kernels/ops/cuda/attention/prefix_shared_attn.py": "rl_engine/backends/cuda/attention/prefix_shared_attn.py", "rl_engine/kernels/ops/cuda/attention/strict_runtime.py": "rl_engine/backends/cuda/attention/strict_runtime.py", "rl_engine/kernels/ops/cuda/ffn.py": "rl_engine/backends/cuda/ffn/ffn.py", + "rl_engine/kernels/ops/cuda/h3/__init__.py": "rl_engine/backends/cuda/model_specific/minimax_h3/__init__.py", + "rl_engine/kernels/ops/cuda/h3/adaln_modulation.py": "rl_engine/backends/cuda/model_specific/minimax_h3/adaln_modulation.py", + "rl_engine/kernels/ops/cuda/h3/adaln_projection.py": "rl_engine/backends/cuda/model_specific/minimax_h3/adaln_projection.py", + "rl_engine/kernels/ops/cuda/h3/adaln_row_gather.py": "rl_engine/backends/cuda/model_specific/minimax_h3/adaln_row_gather.py", + "rl_engine/kernels/ops/cuda/h3/det_linear.py": "rl_engine/backends/cuda/model_specific/minimax_h3/det_linear.py", + "rl_engine/kernels/ops/cuda/h3/final_adaln_out.py": "rl_engine/backends/cuda/model_specific/minimax_h3/final_adaln_out.py", + "rl_engine/kernels/ops/cuda/h3/gate_residual.py": "rl_engine/backends/cuda/model_specific/minimax_h3/gate_residual.py", + "rl_engine/kernels/ops/cuda/h3/rmsnorm.py": "rl_engine/backends/cuda/model_specific/minimax_h3/rmsnorm.py", + "rl_engine/kernels/ops/cuda/h3/sp_norm_adaln.py": "rl_engine/backends/cuda/model_specific/minimax_h3/sp_norm_adaln.py", + "rl_engine/kernels/ops/cuda/h3/timestep_mlp.py": "rl_engine/backends/cuda/model_specific/minimax_h3/timestep_mlp.py", + "rl_engine/kernels/ops/cuda/h3/timestep_sinusoid.py": "rl_engine/backends/cuda/model_specific/minimax_h3/timestep_sinusoid.py", + "rl_engine/kernels/ops/cuda/h3/tp_adaln_projection.py": "rl_engine/backends/cuda/model_specific/minimax_h3/tp_adaln_projection.py", + "rl_engine/kernels/ops/cuda/h3/ws2_comm.py": "rl_engine/backends/cuda/model_specific/minimax_h3/ws2_comm.py", "rl_engine/kernels/ops/cuda/linear/__init__.py": "rl_engine/backends/cuda/embedding/__init__.py", "rl_engine/kernels/ops/cuda/linear/embedding.py": "rl_engine/backends/cuda/embedding/embedding.py", "rl_engine/kernels/ops/cuda/linear/lm_head.py": "rl_engine/backends/cuda/embedding/lm_head.py", @@ -705,6 +753,15 @@ "rl_engine/kernels/ops/pytorch/attention/stateful_kv.py": "rl_engine/reference/attention/stateful_kv.py", "rl_engine/kernels/ops/pytorch/ffn/__init__.py": "rl_engine/ops/ffn/__init__.py", "rl_engine/kernels/ops/pytorch/ffn/ffn.py": "rl_engine/ops/ffn/qwen3.py", + "rl_engine/kernels/ops/pytorch/h3/__init__.py": "rl_engine/reference/minimax_h3/__init__.py", + "rl_engine/kernels/ops/pytorch/h3/adaln_projection.py": "rl_engine/reference/minimax_h3/adaln_projection.py", + "rl_engine/kernels/ops/pytorch/h3/adaln_row_gather.py": "rl_engine/reference/minimax_h3/adaln_row_gather.py", + "rl_engine/kernels/ops/pytorch/h3/final_adaln_out.py": "rl_engine/reference/minimax_h3/final_adaln_out.py", + "rl_engine/kernels/ops/pytorch/h3/fixed_order.py": "rl_engine/reference/minimax_h3/fixed_order.py", + "rl_engine/kernels/ops/pytorch/h3/gate_residual.py": "rl_engine/reference/minimax_h3/gate_residual.py", + "rl_engine/kernels/ops/pytorch/h3/rmsnorm.py": "rl_engine/reference/minimax_h3/rmsnorm.py", + "rl_engine/kernels/ops/pytorch/h3/timestep_mlp.py": "rl_engine/reference/minimax_h3/timestep_mlp.py", + "rl_engine/kernels/ops/pytorch/h3/timestep_sinusoid.py": "rl_engine/reference/minimax_h3/timestep_sinusoid.py", "rl_engine/kernels/ops/pytorch/linear/__init__.py": "rl_engine/reference/embedding/__init__.py", "rl_engine/kernels/ops/pytorch/linear/embedding.py": "rl_engine/reference/embedding/embedding.py", "rl_engine/kernels/ops/pytorch/linear/lm_head.py": "rl_engine/reference/embedding/lm_head.py", @@ -776,6 +833,13 @@ "rl_engine/testing/__init__.py": "rl_engine/validation/common/__init__.py", "rl_engine/testing/attention_comparison.py": "rl_engine/validation/common/attention_comparison.py", "rl_engine/testing/distributed_logprob_comparison.py": "rl_engine/validation/common/distributed_logprob_comparison.py", + "rl_engine/testing/h3_cases.py": "rl_engine/validation/models/h3_cases.py", + "rl_engine/testing/h3_chain.py": "rl_engine/validation/models/h3_chain.py", + "rl_engine/testing/h3_manifest.json": "rl_engine/validation/models/h3_manifest.json", + "rl_engine/testing/h3_provider.py": "rl_engine/validation/models/h3_provider.py", + "rl_engine/testing/h3_report.py": "rl_engine/validation/models/h3_report.py", + "rl_engine/testing/h3_weights.py": "rl_engine/validation/models/h3_weights.py", + "rl_engine/testing/h3_ws2.py": "rl_engine/validation/models/h3_ws2.py", "rl_engine/testing/logprob_comparison.py": "rl_engine/validation/common/logprob_comparison.py", "rl_engine/testing/logprob_drift.py": "rl_engine/validation/common/logprob_drift.py", "rl_engine/testing/reference_ops.py": "rl_engine/reference/definitions.py", @@ -790,6 +854,13 @@ "scripts/check_rocm_env.py": "tools/env/check_rocm_env.py", "scripts/check_stateful_kv.py": "tools/validation/operators/check_stateful_kv.py", "scripts/ci_smoke.py": "tools/checks/ci_smoke.py", + "scripts/h3_chain_replay.py": "tools/validation/models/h3_chain_replay.py", + "scripts/h3_evidence.py": "tools/validation/models/h3_evidence.py", + "scripts/h3_prior_art.py": "tools/validation/models/h3_prior_art.py", + "scripts/h3_ws2_evidence.py": "tools/validation/models/h3_ws2_evidence.py", + "scripts/plot_h3_evidence.py": "tools/validation/models/plot_h3_evidence.py", + "scripts/plot_h3_prior_art.py": "tools/validation/models/plot_h3_prior_art.py", + "scripts/prepare_h3_weights.py": "tools/weights/prepare_h3_weights.py", "scripts/prepare_ws1_weights.py": "tools/weights/prepare_ws1_weights.py", "scripts/run_perf.py": "tools/benchmarking/run_perf.py", "scripts/run_profile_suite.py": "tools/benchmarking/run_profile_suite.py", @@ -820,6 +891,27 @@ "tests/distributed/test_rocm_attention_transport.py": "tests/distributed/cp/test_rocm_attention_transport.py", "tests/distributed/test_rocm_strict_attention_cp.py": "tests/distributed/cp/test_rocm_strict_attention_cp.py", "tests/distributed/test_transport_deterministic_collective.py": "tests/distributed/collectives/test_transport_deterministic_collective.py", + "tests/h3/conftest.py": "tests/models/minimax_h3/conftest.py", + "tests/h3/test_h3_adaln_gate_residual.py": "tests/models/minimax_h3/test_h3_adaln_gate_residual.py", + "tests/h3/test_h3_adaln_modulation.py": "tests/models/minimax_h3/test_h3_adaln_modulation.py", + "tests/h3/test_h3_adaln_projection.py": "tests/models/minimax_h3/test_h3_adaln_projection.py", + "tests/h3/test_h3_adaln_row_gather.py": "tests/models/minimax_h3/test_h3_adaln_row_gather.py", + "tests/h3/test_h3_benchmark_timing.py": "tests/models/minimax_h3/test_h3_benchmark_timing.py", + "tests/h3/test_h3_chain_replay_cli.py": "tests/models/minimax_h3/test_h3_chain_replay_cli.py", + "tests/h3/test_h3_cli.py": "tests/models/minimax_h3/test_h3_cli.py", + "tests/h3/test_h3_conditioning_e2e.py": "tests/models/minimax_h3/test_h3_conditioning_e2e.py", + "tests/h3/test_h3_det_linear.py": "tests/models/minimax_h3/test_h3_det_linear.py", + "tests/h3/test_h3_final_adaln_out.py": "tests/models/minimax_h3/test_h3_final_adaln_out.py", + "tests/h3/test_h3_native_validation.py": "tests/models/minimax_h3/test_h3_native_validation.py", + "tests/h3/test_h3_report.py": "tests/models/minimax_h3/test_h3_report.py", + "tests/h3/test_h3_rmsnorm.py": "tests/models/minimax_h3/test_h3_rmsnorm.py", + "tests/h3/test_h3_sp_norm_adaln.py": "tests/models/minimax_h3/test_h3_sp_norm_adaln.py", + "tests/h3/test_h3_timestep_mlp.py": "tests/models/minimax_h3/test_h3_timestep_mlp.py", + "tests/h3/test_h3_timestep_sinusoid.py": "tests/models/minimax_h3/test_h3_timestep_sinusoid.py", + "tests/h3/test_h3_tp_adaln.py": "tests/models/minimax_h3/test_h3_tp_adaln.py", + "tests/h3/test_h3_weights.py": "tests/models/minimax_h3/test_h3_weights.py", + "tests/h3/test_h3_ws2_processes.py": "tests/models/minimax_h3/test_h3_ws2_processes.py", + "tests/h3/test_prepare_h3_weights.py": "tests/models/minimax_h3/test_prepare_h3_weights.py", "tests/linear_logp_tp.py": "tests/distributed/tp/linear_logp_tp.py", "tests/test_alignment_model_wrappers.py": "tests/models/qwen3/test_alignment_model_wrappers.py", "tests/test_alignment_wrapper_interfaces.py": "tests/models/qwen3/test_alignment_wrapper_interfaces.py", diff --git a/tools/validation/models/h3_block_replay.py b/tools/validation/models/h3_block_replay.py new file mode 100644 index 000000000..7d8230294 --- /dev/null +++ b/tools/validation/models/h3_block_replay.py @@ -0,0 +1,114 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Replay one real MiniMax-H3 block forward and backward and write JSON evidence (RFC #420). + +See ``rl_engine/validation/models/h3_block.py`` for what is compared. + + python tools/weights/prepare_h3_weights.py --out --block + export RL_KERNEL_H3_WEIGHTS= + python tools/validation/models/h3_block_replay.py --out block_replay.json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 + +from rl_engine.validation.models import h3_block, h3_chain # noqa: E402 +from rl_engine.validation.models.h3_cases import H3_BLOCK_LAYOUTS # noqa: E402 +from rl_engine.validation.models.h3_weights import ( # noqa: E402 + load_h3_block_weights, + load_h3_manifest, +) + + +def main() -> None: + """Run the forward (and optionally backward) block replay per layout and save the report.""" + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--layouts", + default="tiny,small,medium", + help=f"comma list of packed FL2VA layouts: {', '.join(H3_BLOCK_LAYOUTS)}", + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--no-backward", action="store_true", help="forward replay only") + parser.add_argument("--timing-repeats", type=int, default=5) + parser.add_argument("--out", type=Path, default=None) + args = parser.parse_args() + + layouts = args.layouts.split(",") + unknown = sorted(set(layouts) - set(H3_BLOCK_LAYOUTS)) + if unknown: + parser.error(f"unknown layouts: {unknown}") + if not torch.cuda.is_available(): + raise SystemExit("the block replay needs a CUDA device") + + weights = load_h3_block_weights("cuda") + manifest = load_h3_manifest() + ops = h3_block.CandidateOps() + + forward, backward = [], [] + for layout in layouts: + case = h3_block.run_block_case( + weights, layout=layout, seed=args.seed, ops=ops, timing_repeats=args.timing_repeats + ) + forward.append(case) + print( + f"{layout} S={case['seq_len']}: first_drift={case['first_drift']} " + f"first_isolated_drift={case['first_isolated_drift']} " + f"repeat_drift={case['first_repeat_drift']} batch_drift={case['first_batch_drift']} " + f"out err cand {case['output']['candidate_rel_err_vs_golden']:.2e} " + f"prov {case['output']['provider_rel_err_vs_golden']:.2e}" + ) + if not args.no_backward: + case = h3_block.run_block_backward( + weights, layout=layout, seed=args.seed, ops=ops, timing_repeats=args.timing_repeats + ) + backward.append(case) + det = all(v["repeat_bitwise_equal"] for v in case["leaves"].values()) + print( + f"{layout} backward: leaves {'det' if det else 'NONDET'}, " + f"grad_repeat_drift={case['first_grad_repeat_drift']} " + f"grad_batch_drift={case['first_grad_batch_drift']} " + f"dX batch-row {case['leaves']['hidden']['batch_row_bitwise_equal']}" + ) + torch.cuda.empty_cache() + + evidence = { + "kind": "h3_one_block_replay", + "rfc": manifest["rfc"], + "model_revision": manifest["model_identity"]["revision"], + "weights_sha256": manifest["weight_shards"], + "reference_commit": manifest["reference_implementation"]["commit"], + **h3_chain.git_state(), + "environment": h3_chain.environment(), + "nodes": [ + { + "node": n.name, + "inputs": list(n.inputs), + "rfc_row": n.binding.rfc_row, + "status": n.binding.status, + "backend": n.binding.backend, + } + for n in h3_block.NODES + ], + "forward_cases": forward, + "backward_cases": backward, + } + if args.out is not None: + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(evidence, indent=2) + "\n") + print(f"wrote {args.out}") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/h3_chain_replay.py b/tools/validation/models/h3_chain_replay.py new file mode 100644 index 000000000..11068259c --- /dev/null +++ b/tools/validation/models/h3_chain_replay.py @@ -0,0 +1,137 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Replay the MiniMax-H3 conditioning chain and write JSON evidence (RFC #420). + +See ``rl_engine/validation/models/h3_chain.py`` for what is compared. + + export RL_KERNEL_H3_WEIGHTS= + python tools/validation/models/h3_chain_replay.py --out chain_replay.json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 + +from rl_engine.runtime.registry import KernelRegistry # noqa: E402 +from rl_engine.validation.models import h3_chain # noqa: E402 +from rl_engine.validation.models.h3_weights import ( # noqa: E402 + load_h3_conditioning_weights, + load_h3_manifest, +) + + +def _ints(text: str) -> list[int]: + """Parse a comma-separated list of integer case sizes.""" + + return [int(v) for v in text.split(",")] + + +def main() -> None: + """Replay requested CUDA stages with prerequisites and optionally save JSON evidence.""" + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--stages", + default=",".join(s.name for s in h3_chain.STAGES), + help="comma list of stages to report; preceding stages run automatically", + ) + parser.add_argument("--timesteps", type=_ints, default=[1, 2, 4], help="comma list of T") + parser.add_argument("--seq-lens", type=_ints, default=[3, 257, 4097], help="comma list of S") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--out", type=Path, default=None) + parser.add_argument( + "--backward", + action="store_true", + help="also replay the chain backward (needs every stage): parameter-gradient " + "determinism and accuracy for the separate-op, fused and diffusers chains", + ) + args = parser.parse_args() + + wanted = args.stages.split(",") + stage_names = [s.name for s in h3_chain.STAGES] + unknown = sorted(set(wanted) - set(stage_names)) + if unknown: + parser.error(f"unknown stages: {unknown}") + stages = [s for s in h3_chain.STAGES if s.name in wanted] + if args.backward and len(stages) != len(h3_chain.STAGES): + parser.error("--backward requires every stage of the chain") + last_stage = max(i for i, stage in enumerate(h3_chain.STAGES) if stage.name in wanted) + execution_stages = h3_chain.STAGES[: last_stage + 1] + + if not torch.cuda.is_available(): + raise SystemExit("the chain replay needs a CUDA device") + + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + weights = load_h3_conditioning_weights("cuda") + registry = KernelRegistry() + manifest = load_h3_manifest() + + cases = [] + for num_timesteps in args.timesteps: + for seq_len in args.seq_lens: + case = h3_chain.run_case( + registry, + weights, + num_timesteps=num_timesteps, + seq_len=seq_len, + seed=args.seed, + stages=execution_stages, + ) + case["stages"] = [entry for entry in case["stages"] if entry["stage"] in wanted] + cases.append(case) + summary = ", ".join( + f"{e['stage']}=" + f"{'eq' if e['chained_vs_provider']['bitwise_equal'] else 'drift'}" + f"(gold {e['chained_vs_golden']['max_abs']:.3e})" + for e in case["stages"] + ) + print(f"T={num_timesteps} S={seq_len}: {summary}; first_drift={case['first_drift']}") + + backward_cases = [] + if args.backward: + for num_timesteps in args.timesteps: + for seq_len in args.seq_lens: + case = h3_chain.run_backward_case( + registry, weights, num_timesteps=num_timesteps, seq_len=seq_len, seed=args.seed + ) + backward_cases.append(case) + parts = [] + for mode in h3_chain.BACKWARD_MODES: + entries = case["leaves"].values() + det = all(e[mode]["repeat_bitwise_equal"] for e in entries) + worst = max(e[mode]["max_abs_vs_golden_over_absmax"] for e in entries) + parts.append(f"{mode}: {'det' if det else 'NONDET'} {worst:.1e}") + summary = " | ".join(parts) + print(f"backward T={num_timesteps} S={seq_len}: {summary}") + + evidence = { + "kind": "h3_conditioning_chain_replay", + "rfc": manifest["rfc"], + "model_revision": manifest["model_identity"]["revision"], + "weights_sha256": manifest["weight_shards"], + "reference_commit": manifest["reference_implementation"]["commit"], + **h3_chain.git_state(), + "environment": h3_chain.environment(), + "stages": [s.name for s in stages], + "executed_stages": [s.name for s in execution_stages], + "cases": cases, + "backward_cases": backward_cases, + } + if args.out is not None: + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(evidence, indent=2) + "\n") + print(f"wrote {args.out}") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/h3_evidence.py b/tools/validation/models/h3_evidence.py new file mode 100644 index 000000000..213d5beb3 --- /dev/null +++ b/tools/validation/models/h3_evidence.py @@ -0,0 +1,74 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Write one RFC #420 operator's evidence report (performance + accuracy) as JSON. + +Run it on a clean tree; the report records the commit and whether the tree +was dirty. ``tools/validation/models/plot_h3_evidence.py`` turns the report into a figure. + + export RL_KERNEL_H3_WEIGHTS= + python tools/validation/models/h3_evidence.py --op timestep_sinusoid_h3 \\ + --out reports/experiments/h3-timestep-sinusoid-b200/report.json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 + +from rl_engine.runtime.registry import KernelRegistry # noqa: E402 +from rl_engine.validation.models.h3_chain import environment, git_state # noqa: E402 +from rl_engine.validation.models.h3_report import ACCURACY, PERF_CASES, measure # noqa: E402 +from rl_engine.validation.models.h3_weights import ( # noqa: E402 + WEIGHTS_ENV, + h3_weights_dir, + load_h3_manifest, +) + + +def main() -> None: + """Collect CUDA accuracy and performance evidence, requiring weights for weighted ops.""" + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--op", required=True, choices=sorted(ACCURACY)) + parser.add_argument("--out", type=Path, required=True) + args = parser.parse_args() + uses_weights = args.op != "timestep_sinusoid_h3" + if uses_weights and h3_weights_dir() is None: + raise SystemExit(f"set {WEIGHTS_ENV} to the tools/weights/prepare_h3_weights.py output dir") + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + + registry = KernelRegistry() + manifest = load_h3_manifest() + report = { + "kind": "h3_operator_report", + "op": args.op, + "rfc": manifest["rfc"], + "model_revision": manifest["model_identity"]["revision"], + "weight_source": "pinned_checkpoint" if uses_weights else "not_applicable", + "reference_commit": manifest["reference_implementation"]["commit"], + **git_state(), + "environment": environment(), + "accuracy": ACCURACY[args.op](registry), + "perf": [measure(case) for case in PERF_CASES[args.op](registry)], + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + print( + f"wrote {args.out} (commit {report['rl_kernel_commit'][:7]}, " + f"dirty={report['tracked_tree_dirty']})" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/h3_prior_art.py b/tools/validation/models/h3_prior_art.py new file mode 100644 index 000000000..9012542e2 --- /dev/null +++ b/tools/validation/models/h3_prior_art.py @@ -0,0 +1,946 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Compare one RFC #420 operator with existing implementations: batch invariance, accuracy, latency. + +The RFC #420 reuse rule asks, for every operator, whether an existing implementation is +batch-invariant (then reuse it) and for accuracy and performance comparisons. This runner +measures the operator's rl-kernel CUDA op next to the existing implementations that import in +the current environment, and writes one JSON report per operator: + + python tools/validation/models/h3_prior_art.py --op adaln_row_gather \\ + --out reports/experiments/h3-prior-art-b200/adaln_row_gather.json \\ + [--megatron-src /path/to/Megatron-LM] [--quick] + python tools/validation/models/plot_h3_prior_art.py \\ + reports/experiments/h3-prior-art-b200/adaln_row_gather.json + +Only operators whose rl-kernel op exists on the current branch can be selected. Optional +libraries (SGLang, Liger, Transformer Engine, vLLM, Megatron-LM from a source tree) are skipped +and recorded as unavailable when they do not import. Global batch-invariant modes (vLLM, +SGLang, Megatron) patch aten for the whole process, so each runs in its own subprocess, with +the environment variables it relies on set before CUDA initialises. + +Batch invariance is bitwise, for the forward, the per-row gradients, and repeatability of every +parameter and table gradient, with three checks: + +1. every row computed alone vs inside full batches of 64, 257 and 2048 rows, seeds 3-5; +2. the full workload-size batch vs sub-batches that together cover every row; +3. eight probe rows at the front, middle and back of batches of every size 1..9 and + 2^k - 1, 2^k, 2^k + 1, plus the whole batch reversed. + +Sparse probes (a few rows, a few batch sizes) miss row-specific or large-batch dependence, which +is why every row is checked. Accuracy is max|err| / max|ref| against the same computation in +FP64; latency is the median of 100 CUDA-event samples after 20 warm-ups. Run on a clean tree on +an otherwise idle GPU; the report records the commit, the tree state and every library version. +""" + +from __future__ import annotations + +import argparse +import contextlib +import importlib.metadata as md +import importlib.util +import json +import os +import statistics +import subprocess +import sys +from pathlib import Path +from typing import Any, Callable + +REPO_ROOT = Path(__file__).resolve().parents[3] +sys.path.insert(0, str(REPO_ROOT)) + +import torch # noqa: E402 +import torch.nn.functional as F # noqa: E402 + +H, D, F1, T3, EPS = 5376, 2688, 256, 3, 1e-5 + +#: operator -> rl-kernel module that must exist on the branch +OPS = { + "timestep_sinusoid": "rl_engine.backends.cuda.model_specific.minimax_h3.timestep_sinusoid", + "timestep_mlp": "rl_engine.backends.cuda.model_specific.minimax_h3.timestep_mlp", + "adaln_projection": "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection", + "adaln_row_gather": "rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather", + "norm_modulate": "rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm", + "gate_residual": "rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual", + "final_adaln_out": "rl_engine.backends.cuda.model_specific.minimax_h3.final_adaln_out", +} +#: global batch-invariant modes compared for the GEMM operators, with their environment +MODES = { + "plain": {}, + "vllm": {"CUBLAS_WORKSPACE_CONFIG": ":16:8", "CUBLASLT_WORKSPACE_SIZE": "1"}, + "sglang": {}, + "sglang_ieee": {"TRITON_F32_DEFAULT": "ieee"}, + "megatron_te_native": {"CUBLASLT_WORKSPACE_SIZE": "0"}, + "megatron_triton": {}, + "megatron_triton_ieee": {"TRITON_F32_DEFAULT": "ieee"}, +} +GEMM_OPS = ("timestep_mlp", "adaln_projection") +LIBS = ("torch", "triton", "diffusers", "vllm", "flashinfer-python", "sglang", "liger-kernel") +LIBS += ("transformer_engine", "megatron-core") + +DEV = "cuda" + + +# --------------------------------------------------------------------------- # +# Batch-invariance checks +# --------------------------------------------------------------------------- # + + +class _Case: + """``fn(**inputs)`` with row arguments sliced together; fixed upstream gradients per seed.""" + + def __init__(self, fn, inputs, rows, grads, out_dims, seed): + self.fn, self.inputs, self.rows, self.grads, self.out_dims = ( + fn, + inputs, + rows, + grads, + out_dims, + ) + self.n = next(inputs[k].shape[d] for k, d in rows.items()) + outs = self._flat(fn(**inputs)) + gen = torch.Generator(device=DEV).manual_seed(seed) + self.ups = [torch.randn(o.shape, device=DEV, generator=gen).to(o.dtype) for o in outs] + + @staticmethod + def _flat(out): + return list(out) if isinstance(out, (tuple, list)) else [out] + + def run(self, idx=None): + args = {} + for k, v in self.inputs.items(): + if idx is not None and k in self.rows: + v = v.index_select(self.rows[k], idx) + if k in self.grads: + v = v.detach().clone().requires_grad_(True) + args[k] = v + outs = self._flat(self.fn(**args)) + ups = ( + self.ups + if idx is None + else [u.index_select(d, idx) for u, d in zip(self.ups, self.out_dims)] + ) + live = [(o, u) for o, u in zip(outs, ups) if o.requires_grad] + if live: + torch.autograd.backward([o for o, _ in live], [u for _, u in live]) + grads = {k: args[k].grad for k in self.grads} + return [o.detach() for o in outs], grads + + +def _sweep_sizes(n: int) -> list[int]: + sizes = set(range(1, min(n, 9) + 1)) + k = 16 + while k <= n: + sizes.update(v for v in (k - 1, k, k + 1) if v <= n) + k *= 2 + return sorted(sizes | {n}) + + +def batch_invariance(cand: dict[str, Any], quick: bool) -> dict[str, Any]: + fn, make, rows, grads = cand["fn"], cand["make"], cand["rows"], cand["grads"] + out_dims = cand.get("out_dims", (0,)) + params = [k for k in grads if k not in rows] + full_sizes, big, subs = cand["bi_sizes"] + if quick: + full_sizes, big, subs = (17, 64), 256, (1, 7, 64) + seeds = (3, 4, 5) + result: dict[str, Any] = {"all_rows": {}, "params_repeatable": None} + ok = True + + def same_row(a, b, d, i): + return torch.equal(a.select(d, 0), b.select(d, i)) + + for n in full_sizes: + fwd = grd = 0 + repeat = True + for seed in seeds: + case = _Case(fn, make(seed, n), rows, grads, out_dims, seed) + fo, fg = case.run() + for i in range(n): + o, g = case.run(torch.tensor([i], device=DEV)) + fwd += not all(same_row(a, b, d, i) for a, b, d in zip(o, fo, out_dims)) + grd += not all( + same_row(g[k], fg[k], rows[k], i) + for k in rows + if k in grads and g[k] is not None + ) + if params: + _, again = case.run() + repeat &= all(torch.equal(again[k], fg[k]) for k in params if fg[k] is not None) + result["all_rows"][str(n)] = { + "rows": len(seeds) * n, + "fwd_differ": fwd, + "rowgrad_differ": grd, + } + if params: + result["params_repeatable"] = repeat and result["params_repeatable"] is not False + ok &= fwd == 0 and grd == 0 and repeat + + fwd = grd = checked = 0 + for seed in seeds[:2]: + case = _Case(fn, make(seed, big), rows, grads, out_dims, seed) + fo, fg = case.run() + for sb in subs: + step = sb if sb >= 1024 else max(sb, big // 512) + for start in range(0, big, step): + idx = torch.arange(start, min(start + sb, big), device=DEV) + o, g = case.run(idx) + checked += 1 + fwd += not all( + torch.equal(a, b.index_select(d, idx)) for a, b, d in zip(o, fo, out_dims) + ) + grd += not all( + torch.equal(g[k], fg[k].index_select(rows[k], idx)) + for k in rows + if k in grads and g[k] is not None + ) + del case, fo, fg + torch.cuda.empty_cache() + result["full_vs_sub_batches"] = { + "full_batch": big, + "sub_batch_sizes": list(subs), + "sub_batches": checked, + "fwd_differ": fwd, + "rowgrad_differ": grd, + } + ok &= fwd == 0 and grd == 0 + + n = full_sizes[-1] + sizes = _sweep_sizes(n) + probes = sorted({0, n - 1, *list(range(n // 8 + 3, n - 1, n // 8))[:6]}) + compared = failed = 0 + first: list[dict[str, int]] = [] + for seed in seeds: + case = _Case(fn, make(seed, n), rows, grads, out_dims, seed) + alone = {r: case.run(torch.tensor([r], device=DEV)) for r in probes} + gen = torch.Generator().manual_seed(seed) + + def compare(idx, pos, r, seed=seed, case=case, alone=alone): + nonlocal compared, failed + o, g = case.run(idx.to(DEV)) + ao, ag = alone[r] + good = all( + torch.equal(a.select(d, pos), b.select(d, 0)) for a, b, d in zip(o, ao, out_dims) + ) + good &= all( + torch.equal(g[k].select(rows[k], pos), ag[k].select(rows[k], 0)) + for k in rows + if k in grads and g[k] is not None + ) + compared += 1 + if not good: + failed += 1 + if len(first) < 8: + first.append( + {"seed": seed, "batch": int(idx.numel()), "position": pos, "row": r} + ) + + for m in sizes: + for r in probes: + others = torch.randperm(n, generator=gen) + others = others[others != r][: m - 1] + for pos in sorted({0, (m - 1) // 2, m - 1}): + compare(torch.cat([others[:pos], torch.tensor([r]), others[pos:]]), pos, r) + for r in probes: + compare(torch.arange(n - 1, -1, -1), n - 1 - r, r) + result["batch_size_sweep"] = { + "batch_sizes": sizes, + "probe_rows": probes, + "comparisons": compared, + "differ": failed, + "first_failures": first, + } + ok &= failed == 0 + result["batch_invariant"] = bool(ok) + return result + + +# --------------------------------------------------------------------------- # +# Accuracy and latency +# --------------------------------------------------------------------------- # + + +def _time_us(fn: Callable[[], Any], warmup: int = 20, iters: int = 100) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + samples = [] + for _ in range(iters): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1e3) + return statistics.median(samples) + + +def _rel(a: torch.Tensor, ref: torch.Tensor) -> float: + return ((a.double() - ref).abs().max() / ref.abs().max().clamp_min(1e-300)).item() + + +def accuracy_latency(cand: dict[str, Any], mode_off, quick: bool) -> dict[str, Any]: + out: dict[str, Any] = {} + for size in cand["perf_sizes"][:2] if quick else cand["perf_sizes"]: + inputs = cand["make"](11, size) + names = list(inputs) + grads = [k for k in cand["grads"] if k in names] + leaves = { + k: (v.detach().clone().requires_grad_(True) if k in grads else v) + for k, v in inputs.items() + } + y = cand["fn"](**leaves) + y = torch.cat([t.reshape(-1) for t in y]) if isinstance(y, (tuple, list)) else y + with mode_off(): + ref_in = { + k: (v.detach().double().requires_grad_(k in grads) if v.is_floating_point() else v) + for k, v in inputs.items() + } + ref = cand["ref"](**ref_in) + ref = torch.cat([t.reshape(-1) for t in ref]) if isinstance(ref, (tuple, list)) else ref + row: dict[str, Any] = {"fwd_rel_err": _rel(y, ref.detach())} + if grads: + gen = torch.Generator(device=DEV).manual_seed(9) + up = torch.randn(y.shape, device=DEV, generator=gen).to(y.dtype) + y.backward(up) + with mode_off(): + ref.backward(up.double()) + row["grad_rel_err"] = {k: _rel(leaves[k].grad, ref_in[k].grad) for k in grads} + row["fwd_us"] = _time_us(lambda inputs=inputs: cand["fn"](**inputs)) + if grads: + + def step(inputs=inputs): + lv = { + k: (v.detach().clone().requires_grad_(True) if k in grads else v) + for k, v in inputs.items() + } + o = cand["fn"](**lv) + outs = list(o) if isinstance(o, (tuple, list)) else [o] + torch.autograd.backward(outs, [torch.ones_like(t) for t in outs]) + + row["fwd_bwd_us"] = _time_us(step) + out[str(size)] = row + del inputs, leaves, y, ref + torch.cuda.empty_cache() + return out + + +# --------------------------------------------------------------------------- # +# Candidates +# --------------------------------------------------------------------------- # + + +def _gen(seed: int, salt: int) -> torch.Generator: + return torch.Generator(device=DEV).manual_seed(seed * 1000 + salt) + + +def _rn(seed, salt, *shape, dtype=torch.float32, scale=1.0): + return (torch.randn(*shape, device=DEV, generator=_gen(seed, salt)) * scale).to(dtype) + + +def _layout(n: int, seed: int): + from rl_engine.validation.models.h3_cases import h3_packed_layout + + ti, tags = h3_packed_layout(n, T3, seed=seed) + return ti.contiguous(), tags.contiguous() + + +def _version(name: str) -> str | None: + try: + return md.version(name) + except md.PackageNotFoundError: + return None + + +T_BI = ((64, 257, 2048), 4096, (1, 7, 1024, 2048)) +S_BI = ((64, 257, 2048), 131072, (1, 7, 4097, 32768, 65536)) +T_PERF = (1, 3, 4, 64, 256, 2048) +S_PERF = (4097, 32768) + + +def _sinusoid_ref(t): + half = F1 // 2 + exponent = -torch.log(torch.tensor(10000.0, dtype=torch.float64, device=DEV)) + freqs = torch.exp(exponent * torch.arange(half, device=DEV, dtype=torch.float64) / half) + e = t.double()[:, None] * freqs[None] + return torch.cat([torch.cos(e), torch.sin(e)], -1) + + +def candidates(op: str, mode: str) -> list[dict[str, Any]]: + """Candidate dicts; a factory that raises (missing library) becomes an unavailable entry.""" + + from rl_engine.validation.models import h3_provider as P + + plain = mode == "plain" + found: list[tuple[str, Callable[[], dict[str, Any]]]] = [] + + if op == "timestep_sinusoid": + base = { + "make": lambda s, n: {"t": torch.rand(n, device=DEV, generator=_gen(s, 1))}, + "rows": {"t": 0}, + "grads": [], + "ref": _sinusoid_ref, + "bi_sizes": ((64, 257, 2048), 8192, (1, 7, 2048, 4097)), + "perf_sizes": T_PERF, + } + + def diffusers(): + return { + **base, + "source": "diffusers get_timestep_embedding (op-for-op replay)", + "fn": lambda t: P.provider_time_proj(t), + } + + def sglang(): + from sglang.kernels.ops.diffusion.modulate.timestep_embedding_jit import ( + timestep_embedding, + ) + + return { + **base, + "source": f"SGLang {_version('sglang')} timestep_embedding", + "fn": lambda t: timestep_embedding( + t, F1, flip_sin_to_cos=True, downscale_freq_shift=0.0 + ), + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.timestep_sinusoid import ( + H3TimestepSinusoidCudaOp, + ) + + op_ = H3TimestepSinusoidCudaOp() + return { + **base, + "source": "rl-kernel H3TimestepSinusoidCudaOp (check_range=False)", + "fn": lambda t: op_.forward(t, check_range=False), + } + + found = [("diffusers", diffusers), ("sglang", sglang), ("rl_kernel", ours)] + + elif op == "timestep_mlp": + + def make(s, n): + return { + "x": _rn(s, 2, n, F1), + "w1": _rn(1, 3, H, F1, scale=F1**-0.5), + "b1": _rn(1, 4, H, scale=0.02), + "w2": _rn(1, 5, D, H, scale=H**-0.5), + "b2": _rn(1, 6, D, scale=0.02), + } + + base = { + "make": make, + "rows": {"x": 0}, + "grads": ["x", "w1", "b1", "w2", "b2"], + "ref": lambda x, w1, b1, w2, b2: F.linear(F.silu(F.linear(x, w1, b1)), w2, b2), + "bi_sizes": T_BI, + "perf_sizes": T_PERF, + } + + def diffusers(): + return { + **base, + "source": f"diffusers TimestepEmbedding, FP32 F.linear [{mode}]", + "fn": lambda x, w1, b1, w2, b2: P.provider_time_embedder(x, w1, b1, w2, b2), + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.timestep_mlp import ( + H3TimestepMLPCudaOp, + ) + + op_ = H3TimestepMLPCudaOp() + return { + **base, + "source": "rl-kernel H3TimestepMLPCudaOp", + "fn": lambda x, w1, b1, w2, b2: op_(x, w1, b1, w2, b2), + } + + found = [(f"diffusers[{mode}]", diffusers)] + ([("rl_kernel", ours)] if plain else []) + + elif op == "adaln_projection": + + def make(s, n): + return { + "temb": _rn(s, 7, n, D), + "w": _rn(1, 8, 18 * H, D, dtype=torch.bfloat16, scale=D**-0.5), + "b": _rn(1, 9, 18 * H, dtype=torch.bfloat16, scale=0.02), + } + + base = { + "make": make, + "rows": {"temb": 0}, + "grads": ["temb", "w", "b"], + "ref": lambda temb, w, b: F.linear(F.silu(temb), w, b), + "bi_sizes": T_BI, + "perf_sizes": T_PERF, + } + + def diffusers(): + return { + **base, + "source": f"diffusers AdaLN projection, BF16 F.linear [{mode}]", + "fn": lambda temb, w, b: F.linear(F.silu(temb).to(w.dtype), w, b), + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_projection import ( + H3AdaLNProjectionCudaOp, + ) + + op_ = H3AdaLNProjectionCudaOp() + return { + **base, + "source": "rl-kernel H3AdaLNProjectionCudaOp", + "fn": lambda temb, w, b: op_.forward_table(temb, w, b), + } + + found = [(f"diffusers[{mode}]", diffusers)] + ([("rl_kernel", ours)] if plain else []) + + elif op == "adaln_row_gather": + + def make(s, n): + ti, tags = _layout(n, s) + return { + "table": _rn(1, 14, 3 * T3, 6 * H, dtype=torch.bfloat16, scale=0.1), + "ti": ti, + "tags": tags, + } + + def ref(table, ti, tags): + return tuple(table.index_select(0, ti * 3 + tags).chunk(6, dim=1)) + + base = { + "make": make, + "rows": {"ti": 0, "tags": 0}, + "grads": ["table"], + "out_dims": (0,) * 6, + "ref": ref, + "bi_sizes": S_BI, + "perf_sizes": S_PERF, + } + + def diffusers(): + return {**base, "source": "diffusers six index_select calls", "fn": ref} + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + H3AdaLNRowGatherCudaOp, + ) + + op_ = H3AdaLNRowGatherCudaOp() + return { + **base, + "source": "rl-kernel H3AdaLNRowGatherCudaOp", + "fn": lambda table, ti, tags: op_(table, ti, tags), + } + + found = [("diffusers", diffusers), ("rl_kernel", ours)] + + elif op == "norm_modulate": + + def make(s, n): + ti, tags = _layout(n, s) + w = (0.5 + torch.rand(H, device=DEV, generator=_gen(1, 11))).bfloat16() + return { + "x": _rn(s, 10, n, H, dtype=torch.bfloat16, scale=2.0), + "w": w, + "sh": _rn(1, 12, 3 * T3, H, dtype=torch.bfloat16, scale=0.1), + "sc": _rn(1, 13, 3 * T3, H, dtype=torch.bfloat16, scale=0.1), + "idx": ti * 3 + tags, + } + + def ref(x, w, sh, sc, idx): + n = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) * w + return n * (1 + sc.index_select(0, idx)) + sh.index_select(0, idx) + + base = { + "make": make, + "rows": {"x": 0, "idx": 0}, + "grads": ["x", "w", "sh", "sc"], + "ref": ref, + "bi_sizes": S_BI, + "perf_sizes": S_PERF, + } + + def per_token(fn): + # Liger / SGLang take per-token shift/scale rows: gather them first (index_select). + return lambda x, w, sh, sc, idx: fn( + x, w, sh.index_select(0, idx), sc.index_select(0, idx) + ) + + def diffusers(): + return { + **base, + "source": "diffusers composition (F.rms_norm + index_select modulation)", + "fn": lambda x, w, sh, sc, idx: P.provider_norm_modulate(x, w, sh, sc, idx), + } + + def liger(): + from liger_kernel.ops.modulated_rms_norm import LigerModulatedRMSNormFunction + + return { + **base, + "source": f"Liger {_version('liger-kernel')} modulated RMSNorm + index_select", + "fn": per_token( + lambda x, w, shr, scr: LigerModulatedRMSNormFunction.apply( + x, w, scr, shr, EPS, 0.0, "llama", False + ) + ), + } + + def liger_rlk_gather(): + # Same Liger kernel, with the rows gathered by rl-kernel's deterministic row gather + # (adaln_row_gather) instead of index_select: separates Liger from the gather. + from liger_kernel.ops.modulated_rms_norm import LigerModulatedRMSNormFunction + + from rl_engine.backends.cuda.model_specific.minimax_h3.adaln_row_gather import ( + H3AdaLNRowGatherCudaOp, + ) + + gather = H3AdaLNRowGatherCudaOp() + + def fn(x, w, sh, sc, idx): + pad = torch.zeros_like(sh) + table = torch.cat([sh, sc, pad, pad, pad, pad], dim=1) + rows = gather(table, torch.div(idx, 3, rounding_mode="floor"), idx % 3) + return LigerModulatedRMSNormFunction.apply( + x, w, rows[1], rows[0], EPS, 0.0, "llama", False + ) + + version = _version("liger-kernel") + return { + **base, + "source": f"Liger {version} modulated RMSNorm + rl-kernel row gather", + "fn": fn, + } + + def sglang(): + from sglang.kernels.ops.diffusion.norm.scale_residual_norm_cutedsl import ( + fused_norm_scale_shift, + ) + + return { + **base, + "grads": [], + "source": f"SGLang {_version('sglang')} fused_norm_scale_shift (forward only)", + "fn": per_token( + lambda x, w, shr, scr: fused_norm_scale_shift( + x[None], w, None, scr[None], shr[None], "rms", EPS + )[0] + ), + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.rmsnorm import H3RMSNormCudaOp + + op_ = H3RMSNormCudaOp() + return { + **base, + "source": "rl-kernel H3RMSNormCudaOp.forward_modulated", + "fn": lambda x, w, sh, sc, idx: op_.forward_modulated(x, w, sh, sc, idx), + } + + def norm_only(make): + def m(s, n): + d = make(s, n) + return {"x": d["x"], "w": d["w"]} + + return m + + def norm_ref(x, w): + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + EPS) * w + + def torch_norm(): + # diffusers' norm without the modulation, to show which half is batch-invariant + return { + **base, + "make": norm_only(make), + "rows": {"x": 0}, + "grads": ["x", "w"], + "ref": norm_ref, + "source": "torch F.rms_norm (no modulation)", + "fn": lambda x, w: F.rms_norm(x, (H,), w, EPS), + } + + def te_norm(): + import transformer_engine.pytorch as te + + mod = te.RMSNorm(H, eps=EPS, params_dtype=torch.bfloat16, device=DEV) + + def fn(x, w): + # TE's backward reads its own parameter, so copy w in (functional_call + # gives a wrong dx); dweight is then TE's parameter gradient, not checked. + with torch.no_grad(): + mod.weight.copy_(w) + return mod(x) + + return { + **base, + "make": norm_only(make), + "rows": {"x": 0}, + "grads": ["x"], + "ref": norm_ref, + "source": f"TE {_version('transformer_engine')} RMSNorm (no modulation)", + "fn": fn, + } + + found = [ + ("diffusers", diffusers), + ("torch_rms_norm", torch_norm), + ("transformer_engine", te_norm), + ("liger", liger), + ("liger_rlk_gather", liger_rlk_gather), + ("sglang", sglang), + ("rl_kernel", ours), + ] + + elif op == "gate_residual": + + def make(s, n): + ti, tags = _layout(n, s) + return { + "res": _rn(s, 15, n, H, dtype=torch.bfloat16), + "y": _rn(s, 16, n, H, dtype=torch.bfloat16), + "gate": _rn(1, 17, 3 * T3, H, dtype=torch.bfloat16, scale=0.1), + "idx": ti * 3 + tags, + } + + def ref(res, y, gate, idx): + return res + gate.index_select(0, idx) * y + + base = { + "make": make, + "rows": {"res": 0, "y": 0, "idx": 0}, + "grads": ["res", "y", "gate"], + "ref": ref, + "bi_sizes": S_BI, + "perf_sizes": S_PERF, + } + + def diffusers(): + return { + **base, + "source": "diffusers residual + gate.index_select(...) * y", + "fn": lambda res, y, gate, idx: P.provider_gate_residual(res, gate, idx, y), + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.gate_residual import ( + H3GateResidualCudaOp, + ) + + op_ = H3GateResidualCudaOp() + return { + **base, + "source": "rl-kernel H3GateResidualCudaOp", + "fn": lambda res, y, gate, idx: op_(res, y, gate, idx), + } + + found = [("diffusers", diffusers), ("rl_kernel", ours)] + + elif op == "final_adaln_out": + + def make(s, n): + ti, _ = _layout(n, s) + return { + "x": _rn(s, 18, n, H, dtype=torch.bfloat16, scale=2.0), + "nw": (0.5 + torch.rand(H, device=DEV, generator=_gen(1, 19))).bfloat16(), + "temb": _rn(1, 20, T3, D), + "w": _rn(1, 21, 2 * H, D, dtype=torch.bfloat16, scale=D**-0.5), + "b": _rn(1, 22, 2 * H, dtype=torch.bfloat16, scale=0.02), + "ti": ti, + } + + def ref(x, nw, temb, w, b, ti): + return P.provider_final_adaln_out(x, nw, temb, w, b, ti) + + base = { + "make": make, + "rows": {"x": 0, "ti": 0}, + "grads": ["x", "nw", "temb", "w", "b"], + "ref": ref, + "bi_sizes": S_BI, + "perf_sizes": S_PERF, + } + + def diffusers(): + return { + **base, + "source": "diffusers MiniMaxH3AdaLayerNormOut (op-for-op replay)", + "fn": ref, + } + + def ours(): + from rl_engine.backends.cuda.model_specific.minimax_h3.final_adaln_out import ( + H3FinalAdaLNOutCudaOp, + ) + + op_ = H3FinalAdaLNOutCudaOp() + return { + **base, + "source": "rl-kernel H3FinalAdaLNOutCudaOp", + "fn": lambda x, nw, temb, w, b, ti: op_(x, nw, temb, w, b, ti), + } + + found = [("diffusers", diffusers), ("rl_kernel", ours)] + + result = [] + for name, factory in found: + try: + result.append({"name": name, **factory()}) + except Exception as exc: # noqa: BLE001 - optional library missing or unusable + result.append({"name": name, "unavailable": f"{type(exc).__name__}: {exc}"[:300]}) + return result + + +# --------------------------------------------------------------------------- # +# Modes and driver +# --------------------------------------------------------------------------- # + + +def enable_mode(mode: str, megatron_src: str | None): + """Switch on a global batch-invariant mode; returns a context manager that turns it off.""" + + base = mode[:-5] if mode.endswith("_ieee") else mode + if base == "plain": + return contextlib.nullcontext + if base == "vllm": + from vllm.model_executor.determinism.batch_invariant import enable_batch_invariant_mode + + enable_batch_invariant_mode() + return contextlib.nullcontext + if base == "sglang": + from sglang.srt.batch_invariant_ops import batch_invariant_ops as mod + + mod.enable_batch_invariant_mode() + kwargs: dict[str, str] = {} + else: + if megatron_src: + sys.path.insert(0, megatron_src) + from megatron.core.transformer.custom_layers import batch_invariant_kernels as mod + + kwargs = {"backend": base.split("_", 1)[1]} + mod.enable_batch_invariant_mode(**kwargs) + if kwargs["backend"] != "triton": + return contextlib.nullcontext + + @contextlib.contextmanager + def off(): + # The Triton matmuls have no FP64 configuration: FP64 references run with the mode off. + mod.disable_batch_invariant_mode() + try: + yield + finally: + mod.enable_batch_invariant_mode(**kwargs) + + return off + + +def run_mode(op: str, mode: str, megatron_src: str | None, quick: bool) -> dict[str, Any]: + try: + mode_off = enable_mode(mode, megatron_src) + except Exception as exc: # noqa: BLE001 + return {"unavailable": f"{type(exc).__name__}: {exc}"[:300]} + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + out: dict[str, Any] = {} + for cand in candidates(op, mode): + name = cand.pop("name") + if "unavailable" in cand: + out[name] = cand + print(f" [{mode}] {name}: unavailable ({cand['unavailable'][:80]})", flush=True) + continue + entry: dict[str, Any] = {"source": cand["source"], "has_backward": bool(cand["grads"])} + entry["batch_invariance"] = batch_invariance(cand, quick) + entry["accuracy_latency"] = accuracy_latency(cand, mode_off, quick) + out[name] = entry + print( + f" [{mode}] {name}: batch_invariant={entry['batch_invariance']['batch_invariant']}", + flush=True, + ) + torch.cuda.empty_cache() + return out + + +def main() -> None: + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument("--op", required=True, choices=sorted(OPS)) + parser.add_argument("--out", type=Path, required=True) + parser.add_argument("--megatron-src", default=None, help="Megatron-LM source tree (optional)") + parser.add_argument("--quick", action="store_true", help="small sizes, for a smoke test only") + parser.add_argument("--mode", default=None, help=argparse.SUPPRESS) + args = parser.parse_args() + if importlib.util.find_spec(OPS[args.op]) is None: + raise SystemExit(f"{args.op}: {OPS[args.op]} is not on this branch") + if not torch.cuda.is_available(): + raise SystemExit("needs a CUDA device") + + if args.mode is not None: # child process: one global mode + args.out.write_text(json.dumps(run_mode(args.op, args.mode, args.megatron_src, args.quick))) + return + + from rl_engine.validation.models.h3_chain import environment, git_state + + modes = list(MODES) if args.op in GEMM_OPS else ["plain"] + results: dict[str, Any] = {} + for mode in modes: + print(f"[{args.op}] mode {mode}", flush=True) + # The part file sits next to the report, not in /tmp: cluster job epilogs can clear a + # user's /tmp while another of their jobs on the same node is still running. + part = args.out.with_name(f".{args.out.stem}.{mode}.part.json") + part.parent.mkdir(parents=True, exist_ok=True) + part.unlink(missing_ok=True) + cmd = [sys.executable, __file__, "--op", args.op, "--out", str(part), "--mode", mode] + if args.megatron_src: + cmd += ["--megatron-src", args.megatron_src] + if args.quick: + cmd.append("--quick") + proc = subprocess.run(cmd, env={**os.environ, **MODES[mode]}) + results[mode] = ( + json.loads(part.read_text()) + if proc.returncode == 0 and part.exists() + else {"unavailable": f"subprocess exited with {proc.returncode}"} + ) + part.unlink(missing_ok=True) + megatron_commit = None + if args.megatron_src: + megatron_commit = ( + subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=args.megatron_src, + capture_output=True, + text=True, + ).stdout.strip() + or None + ) + report = { + "kind": "h3_prior_art", + "rfc": "RL-Align/RL-Kernel#420", + "op": args.op, + **git_state(), + "environment": { + **environment(), + "libraries": {lib: _version(lib) for lib in LIBS}, + "megatron_source_commit": megatron_commit, + "mode_environment": {m: MODES[m] for m in modes}, + }, + "quick": args.quick, + "results": results, + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + print( + f"wrote {args.out} (commit {report['rl_kernel_commit'][:7]}, " + f"dirty={report['tracked_tree_dirty']})" + ) + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/h3_ws2_evidence.py b/tools/validation/models/h3_ws2_evidence.py new file mode 100644 index 000000000..8ef5976fb --- /dev/null +++ b/tools/validation/models/h3_ws2_evidence.py @@ -0,0 +1,128 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Write an RFC #420 WS2 evidence report (byte equality vs WS1, timing, readback). + +Each TP size runs as real NCCL processes, one GPU per rank, with the +deterministic CUDA collective. Run it on a clean tree on one 8-GPU node; +``tools/validation/models/plot_h3_evidence.py`` turns the report into a figure. + + export RL_KERNEL_H3_WEIGHTS= + python tools/validation/models/h3_ws2_evidence.py --op tp_adaln_3mod --worlds 1,2,4,8 \\ + --out reports/experiments/h3-tp-adaln-3mod-b200/report.json + python tools/validation/models/h3_ws2_evidence.py --op sp_norm_adaln --worlds 1,2,4,8 \\ + --out reports/experiments/h3-sp-norm-adaln-b200/report.json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +import torch # noqa: E402 + +from rl_engine.validation.models.h3_chain import environment, git_state # noqa: E402 +from rl_engine.validation.models.h3_weights import ( # noqa: E402 + load_h3_conditioning_weights, + load_h3_manifest, +) +from rl_engine.validation.models.h3_ws2 import ( # noqa: E402 + naive_sp_mismatch, + run_world, + sp_case, + sp_region_sweep, + tp_projection_sweep, +) + +WEIGHT = "transformer_blocks.0.adaln_proj.linear.weight" +BIAS = "transformer_blocks.0.adaln_proj.linear.bias" + + +def _tp_adaln(worlds: list[int], max_t: int) -> dict: + weights = load_h3_conditioning_weights("cpu") + weight, bias = weights[WEIGHT], weights[BIAS] + g = torch.Generator(device="cpu").manual_seed(0) + inputs = { + "temb": torch.randn(max_t, weight.shape[1], generator=g) * 2, + "weight": weight, + "bias": bias, + "grad": torch.randn(max_t, weight.shape[0], generator=g).bfloat16(), + } + runs = [] + for world in worlds: + ranks = run_world(world, tp_projection_sweep, inputs) + runs.append({"tp": world, "ranks": ranks}) + equal = all( + v for r in ranks for e in r["equality"] for k, v in e.items() if k != "num_timesteps" + ) + print(f"tp={world}: byte-equal to WS1 on every rank and T: {equal}") + return {"weights": [WEIGHT, BIAS], "runs": runs} + + +SP_CASES = [ + {"batch": 1, "seq": 4097, "layout": "block", "seed": 1}, + {"batch": 1, "seq": 4097, "layout": "interleaved", "seed": 2}, + {"batch": 2, "seq": 4097, "layout": "block", "seed": 3}, + {"batch": 1, "seq": 32768, "layout": "block", "seed": 4}, + {"batch": 1, "seq": 32768, "layout": "interleaved", "seed": 5}, + {"batch": 1, "seq": 131072, "layout": "block", "seed": 6}, +] + + +def _sp_norm(worlds: list[int], _max_t: int) -> dict: + runs = [] + for world in worlds: + ranks = run_world(world, sp_region_sweep, {}, cases=SP_CASES) + runs.append({"sp": world, "ranks": ranks}) + equal = all(all(c["equal"].values()) for r in ranks for c in r["cases"]) + print(f"sp={world}: byte-equal to WS1 on every rank and case: {equal}") + naive = { + layout: naive_sp_mismatch(sp_case(1, 32768, layout=layout, seed=4), max(worlds)) + for layout in ("block", "interleaved") + } + region = "norm2(residual + gate_msa[row] * y) * (1 + scale_mlp[row]) + shift_mlp[row]" + return { + "region": region, + "hidden": 5376, + "runs": runs, + "naive_sp": {"sp": max(worlds), **naive}, + } + + +OPS = {"tp_adaln_3mod": _tp_adaln, "sp_norm_adaln": _sp_norm} + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--op", required=True, choices=sorted(OPS)) + parser.add_argument("--worlds", default="1,2,4,8") + parser.add_argument("--max-timesteps", type=int, default=4) + parser.add_argument("--out", type=Path, required=True) + args = parser.parse_args() + worlds = [int(w) for w in args.worlds.split(",")] + if torch.cuda.device_count() < max(worlds): + raise SystemExit(f"needs {max(worlds)} GPUs, found {torch.cuda.device_count()}") + + manifest = load_h3_manifest() + report = { + "kind": "h3_ws2_report", + "op": args.op, + "rfc": manifest["rfc"], + "model_revision": manifest["model_identity"]["revision"], + **git_state(), + "environment": {**environment(), "gpus": torch.cuda.device_count()}, + **OPS[args.op](worlds, args.max_timesteps), + } + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(report, indent=2) + "\n") + commit, dirty = report["rl_kernel_commit"][:7], report["tracked_tree_dirty"] + print(f"wrote {args.out} (commit {commit}, dirty={dirty})") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/plot_h3_block.py b/tools/validation/models/plot_h3_block.py new file mode 100644 index 000000000..a87dac54d --- /dev/null +++ b/tools/validation/models/plot_h3_block.py @@ -0,0 +1,172 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Render the RFC #420 one-block replay (``tools/validation/models/h3_block_replay.py``) as a PNG. + +Needs only the report JSON and matplotlib: + + python tools/validation/models/plot_h3_block.py \ + docs/usage/evidence/h3-one-block-b200/report.json + +writes ``figure.png`` next to the report. Panels: the per-node invariance and +provider checks over every layout; the per-node error against the FP32 golden +for the candidate and diffusers; the same for every backward leaf. +""" + +from __future__ import annotations + +import argparse +import json +import math +from pathlib import Path +from typing import Any + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 +from matplotlib.colors import ListedColormap # noqa: E402 + +RL_KERNEL = "#2a78d6" +PROVIDER = "#eb6834" +PASS = "#1baf7a" +FAIL = "#d64545" +NA = "#e4e3dd" +INK = "#2b2b29" +MUTED = "#6f6e69" +GRID = "#e4e3dd" +SURFACE = "#fcfcfb" +STATUS_MARK = {"row": "", "interim": "*", "reference": "†"} + + +def _style(ax, title: str) -> None: + ax.set_facecolor(SURFACE) + ax.set_title(title, loc="left", fontsize=10.5, color=INK, pad=8) + ax.tick_params(colors=MUTED, labelsize=7.5, length=0) + for side in ("top", "right", "left", "bottom"): + ax.spines[side].set_visible(False) + + +def _checks_panel(ax, report: dict[str, Any]) -> None: + """Rows: nodes. Columns: forward/backward checks, ANDed over every layout.""" + + nodes = [n["node"] for n in report["nodes"]] + status = {n["node"]: n["status"] for n in report["nodes"]} + cols = [ + ("repeat", "forward", "repeat_bitwise_equal"), + ("batch row", "forward", "batch_row_bitwise_equal"), + ("= diffusers\n(isolated)", "forward", None), + ("grad\nrepeat", "backward", "grad_repeat_bitwise_equal"), + ("grad\nbatch row", "backward", "grad_batch_row_bitwise_equal"), + ] + grid = [] + for node in nodes: + row = [] + for _label, kind, key in cols: + cases = report["forward_cases" if kind == "forward" else "backward_cases"] + entries = [e for case in cases for e in case["nodes"] if e["node"] == node] + if key is None: + promised = entries[0]["provider_bitwise_isolated_promised"] + ok = all(e["isolated_vs_provider"]["bitwise_equal"] for e in entries) + row.append(2 if not promised and not ok else (1 if ok else 0)) + else: + values = [e[key] for e in entries] + row.append(2 if not values or None in values else (1 if all(values) else 0)) + grid.append(row) + ax.imshow(grid, cmap=ListedColormap([FAIL, PASS, NA]), vmin=0, vmax=2, aspect="auto") + ax.set_xticks(range(len(cols)), [c[0] for c in cols], fontsize=7.5, color=INK) + ax.set_yticks( + range(len(nodes)), [f"{n}{STATUS_MARK[status[n]]}" for n in nodes], fontsize=7.5, color=INK + ) + ax.xaxis.tick_top() + _style(ax, "Bitwise checks, all layouts") + ax.set_xlabel( + "green = holds grey = not promised (reduction differs from diffusers)\n" + "or not applicable (AdaLN tables have no batch axis)\n" + "* interim operator † provider replay (no kernel yet)", + fontsize=7, + color=MUTED, + ) + + +def _err_panel(ax, labels, cand, prov, title: str) -> None: + xs = range(len(labels)) + ax.plot( + xs, [math.log10(v) for v in cand], "o-", color=RL_KERNEL, lw=1.8, ms=5, label="RL-Kernel" + ) + ax.plot( + xs, [math.log10(v) for v in prov], "s--", color=PROVIDER, lw=1.4, ms=4, label="diffusers" + ) + ax.set_xticks(list(xs), labels, rotation=55, ha="right", fontsize=7, color=INK) + lo = math.floor(min(math.log10(v) for v in cand + prov)) + hi = math.ceil(max(math.log10(v) for v in cand + prov)) + ax.set_yticks(range(lo, hi + 1), [f"1e{v}" for v in range(lo, hi + 1)]) + ax.grid(axis="y", color=GRID, linewidth=0.8) + ax.set_ylabel("max|err| / max|golden|", color=MUTED, fontsize=8) + ax.legend(frameon=False, fontsize=7.5, labelcolor=INK, loc="upper left") + _style(ax, title) + + +def plot(report: dict[str, Any]): + fig, axes = plt.subplots( + 1, 3, figsize=(16, 5.6), facecolor=SURFACE, gridspec_kw={"width_ratios": [1.0, 1.6, 1.3]} + ) + case = max(report["forward_cases"], key=lambda c: c["seq_len"]) + back = ( + max(report["backward_cases"], key=lambda c: c["seq_len"]) + if report["backward_cases"] + else None + ) + _checks_panel(axes[0], report) + _err_panel( + axes[1], + [e["node"] for e in case["nodes"]], + [e["candidate_rel_err_vs_golden"] for e in case["nodes"]], + [e["provider_rel_err_vs_golden"] for e in case["nodes"]], + f"Forward error vs FP32 golden, S = {case['seq_len']}", + ) + if back is not None: + leaves = list(back["leaves"]) + _err_panel( + axes[2], + [f"d {name}" for name in leaves], + [back["leaves"][n]["candidate_rel_err_vs_golden"] for n in leaves], + [back["leaves"][n]["provider_rel_err_vs_golden"] for n in leaves], + f"Gradient error vs FP32 golden, S = {back['seq_len']}", + ) + env = report["environment"] + sizes = ", ".join(str(c["seq_len"]) for c in report["forward_cases"]) + fig.suptitle( + "MiniMax-H3 block 0, forward and backward (RFC #420 ws1_one_h3_block)", + x=0.01, + ha="left", + fontsize=12, + color=INK, + fontweight="bold", + ) + fig.text( + 0.01, + 0.905, + f"{env['gpu']} | torch {env['torch']} | layouts S = {sizes} | " + f"first drift vs diffusers: {case['first_drift']} | " + f"rl-kernel {report['rl_kernel_commit'][:7]}", + fontsize=8, + color=MUTED, + ) + fig.tight_layout(rect=(0, 0, 1, 0.89)) + return fig + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("report", type=Path) + args = parser.parse_args() + report = json.loads(args.report.read_text()) + out = args.report.with_name("figure.png") + plot(report).savefig(out, dpi=150, facecolor=SURFACE) + print(f"wrote {out}") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/plot_h3_evidence.py b/tools/validation/models/plot_h3_evidence.py new file mode 100644 index 000000000..467cb5b7d --- /dev/null +++ b/tools/validation/models/plot_h3_evidence.py @@ -0,0 +1,768 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Render an RFC #420 operator or WS2 report (``tools/validation/models/h3_evidence.py``, +``tools/validation/models/h3_ws2_evidence.py``) as a PNG. + +Needs only the report JSON and matplotlib (not torch or a GPU), so it can run +anywhere: + + python tools/validation/models/plot_h3_evidence.py \\ + reports/experiments/h3-timestep-sinusoid-b200/report.json + +writes ``figure.png`` next to the report. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any, Callable + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.patches # noqa: E402 +import matplotlib.pyplot as plt # noqa: E402 + +# Fixed series identity across every H3 figure (categorical slots 1-3). +RL_KERNEL = "#2a78d6" +PROVIDER = "#eb6834" +THIRD = "#1baf7a" +INK = "#2b2b29" +MUTED = "#6f6e69" +GRID = "#e4e3dd" +SURFACE = "#fcfcfb" + + +def _style(ax, title: str, ylabel: str) -> None: + """Apply the shared report axis styling, title, and vertical-axis label.""" + + ax.set_facecolor(SURFACE) + ax.set_title(title, loc="left", fontsize=11, color=INK, pad=10) + ax.set_ylabel(ylabel, color=MUTED, fontsize=9) + ax.tick_params(colors=MUTED, labelsize=8, length=0) + ax.grid(axis="y", color=GRID, linewidth=0.8) + ax.set_axisbelow(True) + for side in ("top", "right", "left"): + ax.spines[side].set_visible(False) + ax.spines["bottom"].set_color(GRID) + + +def _bars(ax, labels: list[str], series: list[tuple], fmt: str, log: bool = False) -> None: + """Draw grouped series with formatted value labels and a legend; a 4th tuple item is a hatch.""" + + n = len(series) + width = 0.8 / n + for i, (name, color, values, *hatch) in enumerate(series): + xs = [x + (i - (n - 1) / 2) * width for x in range(len(labels))] + bars = ax.bar( + xs, + values, + width * 0.92, + color=color if not hatch else SURFACE, + edgecolor=color, + hatch=hatch[0] if hatch else None, + linewidth=1.2, + label=name, + zorder=2, + ) + for bar, value in zip(bars, values, strict=True): + ax.annotate( + fmt.format(value), + (bar.get_x() + bar.get_width() / 2, bar.get_height()), + xytext=(0, 2), + textcoords="offset points", + ha="center", + fontsize=7, + color=INK, + ) + ax.set_xticks(range(len(labels)), labels) + top = max(v for _, _, values, *_ in series for v in values) + if log: + ax.set_yscale("log") + ax.set_ylim(min(v for _, _, values, *_ in series for v in values) / 2, top * 8) + else: + ax.set_ylim(0, top * 1.35) + ax.legend(frameon=False, fontsize=8, labelcolor=INK, loc="upper left", ncol=min(n, 4)) + + +# Exact zeros cannot sit on a log axis; draw them on this floor and label them. +LOG_FLOOR = 1e-9 + + +def _log_values(values: list[float]) -> list[float]: + """Replace nonpositive values with the plotting floor for logarithmic axes.""" + + return [v if v > 0 else LOG_FLOOR for v in values] + + +def _figure(report: dict[str, Any], suptitle: str): + """Create two report axes with a title and GPU, software, and revision provenance.""" + + fig, axes = plt.subplots(1, 2, figsize=(11, 3.8), facecolor=SURFACE) + env = report["environment"] + fig.suptitle(suptitle, x=0.01, ha="left", fontsize=12, color=INK, fontweight="bold") + fig.text( + 0.01, + 0.01, + f"{env['gpu']} · torch {env['torch']} · CUDA {env['cuda']} · " + f"commit {report['rl_kernel_commit'][:7]} · MiniMax-H3@{report['model_revision'][:7]}", + fontsize=7, + color=MUTED, + ) + return fig, axes + + +def plot_sinusoid(report: dict[str, Any]): + """Return a sinusoid report figure showing latency and errors against FP64.""" + + fig, (left, right) = _figure(report, "timestep_sinusoid_h3") + perf = report["perf"] + _style(left, "Latency per call (lower is better)", "µs") + _bars( + left, + [row["case"] for row in perf], + [ + ("RL-Kernel CUDA", RL_KERNEL, [row["candidate_us"] for row in perf]), + ("CUDA + range check", THIRD, [row["candidate_checked_us"] for row in perf]), + ("diffusers", PROVIDER, [row["provider_us"] for row in perf]), + ], + "{:.0f}", + ) + cases = report["accuracy"]["cases"] + xs = [c["num_timesteps"] for c in cases] + _style(right, "Max |error| vs FP64 golden (log)", "max abs error") + cuda = [c["max_abs_vs_fp64"] for c in cases] + provider = [c["provider_max_abs_vs_fp64"] for c in cases] + right.plot(xs, _log_values(cuda), "o-", color=RL_KERNEL, lw=2, ms=8, label="RL-Kernel CUDA") + right.plot(xs, _log_values(provider), "s--", color=PROVIDER, lw=2, ms=5, label="diffusers") + atol = report["accuracy"]["contract_atol"] + right.axhline(atol, color=MUTED, lw=1, ls=":") + right.annotate( + "contract atol 1e-5", + (xs[-1], atol), + xytext=(0, 3), + textcoords="offset points", + fontsize=7, + color=MUTED, + ha="right", + ) + for x, value in zip(xs, cuda, strict=True): + if value == 0: + right.annotate( + "exact", + (x, LOG_FLOOR), + xytext=(0, 8), + textcoords="offset points", + fontsize=7, + color=INK, + ha="center", + ) + right.set_xscale("log") + right.set_yscale("log") + right.set_ylim(LOG_FLOOR / 3, atol * 30) + right.set_xlabel("number of timesteps T", color=MUTED, fontsize=9) + equal = sum(c["bitwise_equal_to_provider"] for c in cases) + right.legend( + frameon=False, + fontsize=8, + labelcolor=INK, + loc="upper left", + ncol=2, + title=f"bitwise equal to diffusers at {equal}/{len(cases)} sizes", + title_fontsize=8, + alignment="left", + ) + return fig + + +def _latency_panel(ax, perf: list[dict[str, Any]], provider_name: str) -> None: + """Plot candidate and provider latency summaries in microseconds for each case.""" + + _style(ax, "Latency per call (lower is better)", "µs") + _bars( + ax, + [row["case"] for row in perf], + [ + ("RL-Kernel CUDA", RL_KERNEL, [row["candidate_us"] for row in perf]), + (provider_name, PROVIDER, [row["provider_us"] for row in perf]), + ], + "{:.0f}", + ) + + +def _error_boxes(ax, title: str, series: list[tuple[str, str, list[float]]], contract: float): + """Per-draw max |error| as boxes with every draw overlaid, log scale.""" + + _style(ax, title, "max abs error per draw") + for i, (_name, color, values) in enumerate(series): + ax.boxplot( + [values], + positions=[i], + widths=0.45, + showfliers=False, + patch_artist=True, + boxprops={"facecolor": color, "alpha": 0.25, "edgecolor": color}, + medianprops={"color": color, "linewidth": 2}, + whiskerprops={"color": color}, + capprops={"color": color}, + ) + jitter = [i + ((k * 0.6180339) % 1 - 0.5) * 0.3 for k in range(len(values))] + ax.scatter(jitter, values, s=8, color=color, alpha=0.6, linewidths=0, zorder=3) + median = sorted(values)[len(values) // 2] + ax.annotate( + f"median {median:.1e}", + (i + 0.28, median), + fontsize=8, + color=INK, + va="center", + ) + ax.axhline(contract, color=MUTED, lw=1, ls=":") + ax.annotate( + f"contract atol {contract:g}", + (len(series) - 0.5, contract), + xytext=(0, 3), + textcoords="offset points", + fontsize=7, + color=MUTED, + ha="right", + ) + ax.set_yscale("log") + ax.set_xticks(range(len(series)), [name for name, _, _ in series]) + ax.set_xlim(-0.6, len(series) - 0.1) + + +def plot_mlp(report: dict[str, Any]): + """Return an MLP report figure showing latency and per-draw FP64 error distributions.""" + + fig, (left, right) = _figure(report, "timestep_mlp_fp32 (pinned time_embedder weights)") + _latency_panel(left, report["perf"], "diffusers (cuBLAS)") + acc = report["accuracy"] + _error_boxes( + right, + f"Error vs FP64 golden, {acc['draws']} draws of T={acc['num_timesteps']}", + [ + ("RL-Kernel CUDA", RL_KERNEL, acc["cuda_max_abs_vs_fp64"]), + ("diffusers (cuBLAS)", PROVIDER, acc["provider_max_abs_vs_fp64"]), + ], + acc["contract_atol"], + ) + return fig + + +def plot_projection(report: dict[str, Any]): + """Return an AdaLN figure with latency, BF16 rounding accuracy, and the early-cast probe.""" + + fig, (left, right) = _figure(report, "adaln_projection_3mod (pinned block-0 AdaLN weights)") + _latency_panel(left, report["perf"], "diffusers (cuBLAS)") + acc = report["accuracy"] + _style( + right, + f"Outputs equal to the correctly rounded FP64 golden, {acc['draws']} draws", + "% of BF16 outputs", + ) + series = [ + ("RL-Kernel CUDA", RL_KERNEL, [100 * v for v in acc["cuda_correctly_rounded"]]), + ("diffusers (cuBLAS)", PROVIDER, [100 * v for v in acc["provider_correctly_rounded"]]), + ] + for i, (_name, color, values) in enumerate(series): + jitter = [i + ((k * 0.6180339) % 1 - 0.5) * 0.3 for k in range(len(values))] + ax_vals = sorted(values) + median = ax_vals[len(ax_vals) // 2] + right.scatter(jitter, values, s=36, color=color, alpha=0.8, linewidths=0, zorder=3) + right.hlines(median, i - 0.25, i + 0.25, color=color, lw=2.5, zorder=4) + right.annotate( + f"median {median:.3f}%", (i + 0.28, median), fontsize=8, color=INK, va="center" + ) + right.set_xticks(range(len(series)), [name for name, _, _ in series]) + right.set_xlim(-0.6, len(series) - 0.1) + early = sorted(acc["early_cast_golden_match"])[len(acc["early_cast_golden_match"]) // 2] + right.text( + 0.98, + 0.62, + f"an early BF16 cast (probe H7)\nwould match only {100 * early:.0f}%", + transform=right.transAxes, + fontsize=8, + color=MUTED, + ha="right", + ) + return fig + + +def plot_gather(report: dict[str, Any]): + fig, (left, right) = _figure(report, "adaln_row_gather (T = 3, H = 5376, BF16)") + perf = report["perf"] + _style(left, "Forward and backward time (log, lower is better)", "ms") + _bars( + left, + [row["case"] for row in perf], + [ + ("CUDA fwd", RL_KERNEL, [row["candidate_us"] / 1e3 for row in perf]), + ("index_select fwd", PROVIDER, [row["provider_us"] / 1e3 for row in perf]), + ("CUDA bwd", RL_KERNEL, [row["candidate_backward_us"] / 1e3 for row in perf], "////"), + ( + "index_select bwd", + PROVIDER, + [row["provider_backward_us"] / 1e3 for row in perf], + "////", + ), + ], + "{:.2f}", + log=True, + ) + chain = report["accuracy"]["chain_backward"] + _style( + right, + "Whole-chain grads of the FP32 time embedder vs FP64", + "max abs error / golden max", + ) + if not chain: + right.text( + 0.5, + 0.5, + report["accuracy"].get("chain_backward_skipped", "No chain measurements available"), + transform=right.transAxes, + ha="center", + va="center", + color=INK, + fontsize=9, + wrap=True, + ) + right.set_axis_off() + return fig + names = { + "candidate": ("RL-Kernel, separate ops", THIRD, "o"), + "candidate_fused": ("RL-Kernel, fused modulation", RL_KERNEL, "D"), + "provider": ("diffusers", PROVIDER, "s"), + } + labels = [f"T={c['num_timesteps']}\nS={c['seq_len']}" for c in chain] + for j, (mode, (name, color, marker)) in enumerate(names.items()): + leaves = [e for c in chain for e in c["leaves"].values()] + det = sum(e[mode]["repeat_bitwise_equal"] for e in leaves) + adaln = min( + e[mode]["correctly_rounded_fraction"] + for c in chain + for n, e in c["leaves"].items() + if n.startswith("transformer_blocks") + ) + label = ( + f"{name}\n repeat-bitwise {det}/{len(leaves)} · " + f"AdaLN BF16 grads ≥{100 * adaln:.1f}% correctly rounded" + ) + for i, case in enumerate(chain): + values = [ + e[mode]["max_abs_vs_golden_over_absmax"] + for n, e in case["leaves"].items() + if n.startswith("time_embedder") + ] + right.scatter( + [i + (j - 1) * 0.22] * len(values), + values, + s=30, + color=color, + marker=marker, + alpha=0.85, + linewidths=0, + zorder=3, + label=label if i == 0 else None, + ) + right.set_yscale("log") + right.set_ylim(1e-8, 1e4) + right.set_xticks(range(len(labels)), labels, fontsize=7) + right.legend(frameon=False, fontsize=7, labelcolor=INK, loc="upper left") + return fig + + +def plot_rmsnorm(report: dict[str, Any]): + fig, (left, right) = _figure( + report, "h3_rmsnorm (block norm1 + MSA modulation, H = 5376, BF16)" + ) + perf = report["perf"] + _style(left, "Forward and backward time (log, lower is better)", "ms") + _bars( + left, + [row["case"] for row in perf], + [ + ("CUDA fwd", RL_KERNEL, [row["candidate_us"] / 1e3 for row in perf]), + ("diffusers fwd", PROVIDER, [row["provider_us"] / 1e3 for row in perf]), + ("CUDA bwd", RL_KERNEL, [row["candidate_backward_us"] / 1e3 for row in perf], "////"), + ( + "diffusers bwd", + PROVIDER, + [row["provider_backward_us"] / 1e3 for row in perf], + "////", + ), + ], + "{:.2f}", + log=True, + ) + acc = report["accuracy"] + bwd = acc["backward"] + keys = ["dx", "dweight", "dshift", "dscale"] + fwd_ok = acc["modulated_bitwise_vs_diffusers"] and all( + acc["plain_bitwise_vs_nn_rmsnorm"].values() + ) + _style(right, "Backward vs FP64 golden (log, lower is better)", "max abs error / golden max") + + def repeat(mode: str) -> str: + return "yes" if bwd[mode]["repeat_bitwise_equal"] else "no" + + _bars( + right, + keys, + [ + ( + f"RL-Kernel CUDA (repeat-bitwise: {repeat('cuda')})", + RL_KERNEL, + [bwd["cuda"]["rel_error"][k] for k in keys], + ), + ( + f"diffusers (repeat-bitwise: {repeat('provider')})", + PROVIDER, + [bwd["provider"]["rel_error"][k] for k in keys], + ), + ], + "{:.1e}", + log=True, + ) + right.text( + 0.98, + 0.80, + f"forward bitwise equal to diffusers: {'yes' if fwd_ok else 'NO'}", + transform=right.transAxes, + ha="right", + fontsize=8, + color=INK, + ) + return fig + + +def _latency_fwd_bwd(ax, perf: list[dict[str, Any]], provider: str) -> None: + combined = all(row.get("backward_timing_scope") == "forward_and_backward" for row in perf) + backward_label = "fwd+bwd" if combined else "bwd" + _style(ax, "Forward and backward time (log, lower is better)", "ms") + _bars( + ax, + [row["case"] for row in perf], + [ + ("CUDA fwd", RL_KERNEL, [row["candidate_us"] / 1e3 for row in perf]), + (f"{provider} fwd", PROVIDER, [row["provider_us"] / 1e3 for row in perf]), + ( + f"CUDA {backward_label}", + RL_KERNEL, + [row["candidate_backward_us"] / 1e3 for row in perf], + "////", + ), + ( + f"{provider} {backward_label}", + PROVIDER, + [row["provider_backward_us"] / 1e3 for row in perf], + "////", + ), + ], + "{:.2f}", + log=True, + ) + + ax.legend(frameon=False, fontsize=8, labelcolor=INK, loc="upper left", ncol=2) + + +def _backward_errors(ax, backward: dict[str, Any], keys: list[str], note: str) -> None: + _style(ax, "Backward vs FP64 golden (log, lower is better)", "max abs error / golden max") + + def label(name: str, mode: str) -> str: + return ( + f"{name} (repeat-bitwise: {'yes' if backward[mode]['repeat_bitwise_equal'] else 'no'})" + ) + + floor = 1e-9 + _bars( + ax, + keys, + [ + ( + label("RL-Kernel CUDA", "cuda"), + RL_KERNEL, + [max(backward["cuda"]["rel_error"][k], floor) for k in keys], + ), + ( + label("diffusers", "provider"), + PROVIDER, + [max(backward["provider"]["rel_error"][k], floor) for k in keys], + ), + ], + "{:.1e}", + log=True, + ) + for text in ax.texts: + if text.get_text() == f"{floor:.1e}": + text.set_text("exact") + ax.legend(frameon=False, fontsize=8, labelcolor=INK, loc="upper left", ncol=1) + ax.set_ylim(floor / 3, 1e4) + ax.text(0.98, 0.80, note, transform=ax.transAxes, ha="right", fontsize=8, color=INK) + + +def plot_gate_residual(report: dict[str, Any]): + fig, (left, right) = _figure( + report, "adaln_gate_residual (residual + gate_msa[row] * y, H = 5376)" + ) + _latency_fwd_bwd(left, report["perf"], "diffusers") + acc = report["accuracy"] + ok = all(acc["forward_bitwise_vs_diffusers"].values()) + _backward_errors( + right, + acc["backward"], + ["d_residual", "d_sublayer", "d_gate"], + f"forward bitwise equal to diffusers (bf16/fp16/fp32): {'yes' if ok else 'NO'}", + ) + return fig + + +def plot_final(report: dict[str, Any]): + fig, (left, right) = _figure( + report, "final_adaln_out (norm_out: projection + final norm, T = 3)" + ) + _latency_fwd_bwd(left, report["perf"], "diffusers") + acc = report["accuracy"] + _backward_errors( + right, + acc["backward"], + ["dx", "d_norm_w", "d_temb", "dW", "db"], + f"forward equal to diffusers on {100 * acc['equal_to_diffusers_fraction']:.2f}% " + "of elements (projection 1-ULP ties)", + ) + return fig + + +TP_TENSORS = ( + ("table", "table (6 x 3T rows)"), + ("d_temb", "d_temb"), + ("d_weight_shard", "dW shard"), + ("d_bias_shard", "db shard"), +) + + +def plot_tp_adaln(report: dict[str, Any]): + fig, (left, right) = _figure( + report, "tp_adaln_3mod (column-parallel AdaLN projection, 2688 -> 96768, NCCL)" + ) + runs = report["runs"] + _equality_grid( + left, + "Byte-equal to WS1", + [(f"TP{run['tp']}", [e for r in run["ranks"] for e in r["equality"]]) for run in runs], + TP_TENSORS, + "ranks x T", + ) + # Right: slowest rank's time at the largest T; WS1 on one GPU as reference lines. + num_t = max(p["num_timesteps"] for p in runs[0]["ranks"][0]["perf"]) + + def worst(run, key): + perf = [p for r in run["ranks"] for p in r["perf"] if p["num_timesteps"] == num_t] + return max(p[key] for p in perf) / 1e3 + + tp_runs = [run for run in runs if run["tp"] > 1] + _style(right, f"Time per call, T = {num_t} (slowest rank; lower is better)", "ms") + _bars( + right, + [f"TP{run['tp']}" for run in tp_runs], + [ + ("TP fwd", RL_KERNEL, [worst(run, "tp_forward_us") for run in tp_runs]), + ( + "TP fwd+bwd", + RL_KERNEL, + [worst(run, "tp_forward_backward_us") for run in tp_runs], + "////", + ), + ], + "{:.2f}", + log=True, + ) + for key, label, style in ( + ("ws1_forward_us", "WS1 fwd", ":"), + ("ws1_forward_backward_us", "WS1 fwd+bwd", "--"), + ): + value = worst(runs[0], key) + right.axhline( + value, + color=PROVIDER, + linestyle=style, + linewidth=1.4, + zorder=1, + label=f"{label}, 1 GPU: {value:.2f}", + ) + right.legend(frameon=False, fontsize=8, labelcolor=INK, loc="upper left", ncol=2) + return fig + + +def _equality_grid(ax, title, rows, tensors, unit) -> None: + """One cell per (configuration, tensor): how many checks were byte-equal to WS1.""" + + ax.set_facecolor(SURFACE) + ax.set_title(f"{title}; checks = {unit}", loc="left", fontsize=11, color=INK, pad=10) + for y, (_, checks) in enumerate(rows): + for x, (key, _) in enumerate(tensors): + ok = sum(bool(e[key]) for e in checks) + good = ok == len(checks) + ax.add_patch( + matplotlib.patches.FancyBboxPatch( + (x + 0.06, y + 0.08), + 0.88, + 0.84, + boxstyle="round,pad=0,rounding_size=0.06", + facecolor=RL_KERNEL if good else PROVIDER, + edgecolor=SURFACE, + linewidth=2, + ) + ) + label = f"{'equal' if good else 'DIFF'}\n{ok}/{len(checks)}" + ax.text(x + 0.5, y + 0.5, label, ha="center", va="center", fontsize=8, color="white") + ax.set_xlim(0, len(tensors)) + ax.set_ylim(len(rows), 0) + ax.set_xticks([i + 0.5 for i in range(len(tensors))], [label for _, label in tensors]) + ax.set_yticks([i + 0.5 for i in range(len(rows))], [name for name, _ in rows]) + ax.tick_params(colors=MUTED, labelsize=8, length=0) + for side in ax.spines.values(): + side.set_visible(False) + + +SP_TENSORS = ( + ("out", "output rows"), + ("d_residual", "d_residual"), + ("d_y", "d_sublayer"), + ("d_weight", "d_norm_w"), + ("d_table", "d_table"), +) + + +def plot_sp_norm(report: dict[str, Any]): + fig, axes = plt.subplots( + 1, 3, figsize=(15, 3.9), facecolor=SURFACE, gridspec_kw={"width_ratios": [1.35, 1, 0.9]} + ) + env = report["environment"] + fig.suptitle( + "sp_norm_adaln (sequence-parallel norm2(residual + gate * y) with AdaLN modulation, " + "H = 5376, NCCL)", + x=0.01, + ha="left", + fontsize=12, + color=INK, + fontweight="bold", + ) + fig.text( + 0.01, + 0.01, + f"{env['gpu']} x {env['gpus']} · torch {env['torch']} · CUDA {env['cuda']} · " + f"commit {report['rl_kernel_commit'][:7]} · bf16", + fontsize=7, + color=MUTED, + ) + runs = report["runs"] + _equality_grid( + axes[0], + "Byte-equal to WS1", + [ + (f"SP{run['sp']}", [c["equal"] for r in run["ranks"] for c in r["cases"]]) + for run in runs + ], + SP_TENSORS, + "ranks x cases", + ) + + # Middle: the largest block-layout case, slowest rank. + big = max( + (c for c in runs[0]["ranks"][0]["cases"] if c["layout"] == "block"), + key=lambda c: c["seq"] * c["batch"], + ) + + def worst(run, key): + return ( + max( + c[key] + for r in run["ranks"] + for c in r["cases"] + if (c["seq"], c["batch"], c["layout"]) == (big["seq"], big["batch"], big["layout"]) + ) + / 1e3 + ) + + sp_runs = [run for run in runs if run["sp"] > 1] + ax = axes[1] + _style(ax, f"Region fwd+bwd, S = {big['seq']}, block layout", "ms (slowest rank)") + _bars( + ax, + [f"SP{run['sp']}" for run in sp_runs], + [("SP", RL_KERNEL, [worst(run, "sp_forward_backward_us") for run in sp_runs])], + "{:.2f}", + ) + ws1 = worst(runs[0], "ws1_forward_backward_us") + ax.axhline( + ws1, color=PROVIDER, linestyle="--", linewidth=1.4, zorder=1, label=f"WS1, 1 GPU: {ws1:.2f}" + ) + ax.set_ylim(0, ws1 * 1.35) + ax.legend(frameon=False, fontsize=8, labelcolor=INK, loc="upper right") + + # Right: what a naive SP reduction would do instead. + naive = report["naive_sp"] + ax = axes[2] + _style(ax, f"Naive SP{naive['sp']} vs WS1 (S = 32768)", "% of elements that differ") + _bars( + ax, + ["d_norm_w", "d_table"], + [ + ("naive, block", PROVIDER, [100 * naive["block"][k] for k in ("d_weight", "d_table")]), + ( + "naive, interleaved", + PROVIDER, + [100 * naive["interleaved"][k] for k in ("d_weight", "d_table")], + "////", + ), + ("RL-Kernel SP", RL_KERNEL, [0.0, 0.0]), + ], + "{:.1f}", + ) + ax.text( + 0.98, + 0.62, + "naive: each rank's WS1 backward,\nthen a rank-order sum", + transform=ax.transAxes, + ha="right", + fontsize=7, + color=MUTED, + ) + return fig + + +PLOTS: dict[str, Callable[[dict[str, Any]], Any]] = { + "timestep_sinusoid_h3": plot_sinusoid, + "timestep_mlp_fp32": plot_mlp, + "adaln_projection_3mod": plot_projection, + "adaln_row_gather": plot_gather, + "h3_rmsnorm": plot_rmsnorm, + "adaln_gate_residual": plot_gate_residual, + "final_adaln_out": plot_final, + "tp_adaln_3mod": plot_tp_adaln, + "sp_norm_adaln": plot_sp_norm, +} + + +def main() -> None: + """Read an operator report and save its figure as a PNG beside it or at ``--out``.""" + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("report", type=Path) + parser.add_argument("--out", type=Path, default=None) + args = parser.parse_args() + report = json.loads(args.report.read_text()) + fig = PLOTS[report["op"]](report) + fig.tight_layout(rect=(0, 0.04, 1, 0.95)) + out = args.out or args.report.with_name("figure.png") + fig.savefig(out, dpi=150, facecolor=SURFACE) + print(f"wrote {out}") + + +if __name__ == "__main__": + main() diff --git a/tools/validation/models/plot_h3_prior_art.py b/tools/validation/models/plot_h3_prior_art.py new file mode 100644 index 000000000..893612e87 --- /dev/null +++ b/tools/validation/models/plot_h3_prior_art.py @@ -0,0 +1,139 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Plot H3 prior-art latency, accuracy, and batch invariance. + + python tools/validation/models/plot_h3_prior_art.py \\ + reports/experiments/h3-prior-art-b200/.json + +Writes ``.png`` next to the report. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt # noqa: E402 +from matplotlib import transforms # noqa: E402 + + +def _entries(report): + for mode, cands in report["results"].items(): + if "unavailable" in cands: + continue + for name, entry in cands.items(): + if "unavailable" not in entry: + yield (name if name.endswith("]") or mode == "plain" else f"{name}[{mode}]"), entry + + +def _differ(bi) -> int: + rows = sum(v["fwd_differ"] + v["rowgrad_differ"] for v in bi["all_rows"].values()) + sub = bi["full_vs_sub_batches"] + return rows + sub["fwd_differ"] + sub["rowgrad_differ"] + bi["batch_size_sweep"]["differ"] + + +def main() -> None: + path = Path(sys.argv[1]) + report = json.loads(path.read_text()) + entries = list(_entries(report)) + names = [n for n, _ in entries] + fig, axes = plt.subplots(1, 3, figsize=(20, 6.5), layout="constrained") + fig.suptitle( + f"{report['op']} vs existing implementations — {report['environment']['gpu']}, " + f"commit {report['rl_kernel_commit'][:7]}", + fontsize=13, + ) + + ax = axes[0] + for name, entry in entries: + perf = entry["accuracy_latency"] + key = "fwd_bwd_us" if entry["has_backward"] else "fwd_us" + sizes = sorted(perf, key=int) + ax.plot( + [int(s) for s in sizes], + [perf[s][key] for s in sizes], + marker="o", + ls="-" if entry["has_backward"] else "--", + label=name + ("" if entry["has_backward"] else " (forward only)"), + ) + ax.set_xscale("log", base=2) + ax.set_yscale("log") + ax.set_xlabel("rows (timesteps or tokens)") + ax.set_ylabel("µs, median") + ax.set_title("forward + backward latency (dashed: forward only)") + ax.grid(True, which="both", alpha=0.3) + ax.legend(fontsize=7) + + ax = axes[1] + first = entries[0][1]["accuracy_latency"] + size = "3" if "3" in first else sorted(first, key=int)[-1] + ys = range(len(entries)) + fwd = [e["accuracy_latency"][size]["fwd_rel_err"] for _, e in entries] + grad = [ + max(e["accuracy_latency"][size].get("grad_rel_err", {}).values(), default=0) + for _, e in entries + ] + ax.barh([y - 0.2 for y in ys], fwd, 0.4, label="forward") + ax.barh([y + 0.2 for y in ys], [g or float("nan") for g in grad], 0.4, label="worst gradient") + values = [v for v in fwd + grad if v] + if values and max(values) / min(values) > 10: + ax.set_xscale("log") + inside = transforms.blended_transform_factory(ax.transAxes, ax.transData) + for y, a, b in zip(ys, fwd, grad): + ax.text( + 0.98, + y, + f"fwd {a:.1e} / grad {b:.1e}" if b else f"fwd {a:.1e}", + transform=inside, + ha="right", + va="center", + fontsize=7, + bbox={"facecolor": "white", "edgecolor": "none", "alpha": 0.8, "pad": 1}, + ) + ax.set_yticks(list(ys), names, fontsize=7) + ax.invert_yaxis() + ax.set_title(f"error vs FP64, max|err| / max|ref| (size {size})") + ax.legend(fontsize=7, loc="lower right") + ax.grid(True, axis="x", alpha=0.3) + + ax = axes[2] + for y, (_, e) in zip(ys, entries): + bi = e["batch_invariance"] + rep = bi["params_repeatable"] + lines = ["batch-invariant" if bi["batch_invariant"] else "NOT batch-invariant"] + lines.append(f"{_differ(bi)} row comparisons differ") + if rep is not None: + lines.append("parameter/table grads " + ("repeatable" if rep else "NOT repeatable")) + ax.text( + 0.02, + y, + " | ".join(lines), + va="center", + fontsize=8, + color="#3aa676" if bi["batch_invariant"] else "#d14b4b", + ) + ax.set_ylim(len(entries) - 0.5, -0.5) + ax.set_xlim(0, 1) + ax.axis("off") + ax.set_title("batch invariance (bitwise)") + ax.text( + 0.02, + len(entries) - 0.3, + "every row alone vs full batches + full batch vs covering sub-batches + size sweep", + fontsize=7, + color="#555555", + ) + + out = path.with_suffix(".png") + fig.savefig(out, dpi=120) + print(f"wrote {out}") + + +if __name__ == "__main__": + main() diff --git a/tools/weights/prepare_h3_weights.py b/tools/weights/prepare_h3_weights.py new file mode 100644 index 000000000..f6e32bcb5 --- /dev/null +++ b/tools/weights/prepare_h3_weights.py @@ -0,0 +1,125 @@ +#!/usr/bin/env python +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2026 RL-Kernel Contributors + +"""Fetch and pin the MiniMax-H3 conditioning tensors used by the RFC #420 tests. + +Downloads only the config, the shard index and the shards that hold the +manifest tensors (shard 1 of 14, ~4.8 GB) at the pinned revision, verifies +every sha256 against ``rl_engine/validation/models/h3_manifest.json`` and writes the +tensors to ``/h3_conditioning.safetensors``. Point the tests at it with +``export RL_KERNEL_H3_WEIGHTS=``. ``--block`` also writes block 0's +attention/FFN weights (manifest ``block_tensors``, same shard) to +``/h3_block0.safetensors`` for the ``ws1_one_h3_block`` replay. + + python tools/weights/prepare_h3_weights.py --out ~/.cache/rl-kernel/h3 [--block] +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(REPO_ROOT)) + +from rl_engine.validation.models.h3_weights import ( # noqa: E402 + BLOCK_FILE, + EXTRACTED_FILE, + load_h3_manifest, + sha256_file, + sha256_tensor, +) + + +def _check(path: Path, expected: str) -> None: + """Verify a downloaded file's SHA-256, exiting on mismatch or printing success.""" + + actual = sha256_file(path) + if actual != expected: + raise SystemExit(f"sha256 mismatch for {path.name}: {actual} != {expected}") + print(f"ok {path.name} {actual}") + + +def main() -> None: + """Download pinned model artifacts, verify tensor identities, and write extracted weights.""" + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--out", type=Path, required=True, help="output directory") + parser.add_argument( + "--download-dir", + type=Path, + default=None, + help="where the raw shard goes (default: /hf)", + ) + parser.add_argument( + "--block", action="store_true", help="also extract block 0 (manifest block_tensors)" + ) + args = parser.parse_args() + + from huggingface_hub import hf_hub_download + from safetensors import safe_open + from safetensors.torch import save_file + + manifest = load_h3_manifest() + identity = manifest["model_identity"] + out_dir: Path = args.out.expanduser() + download_dir = (args.download_dir or out_dir / "hf").expanduser() + out_dir.mkdir(parents=True, exist_ok=True) + + def fetch(filename: str) -> Path: + """Download one file from the manifest's model subfolder and pinned revision.""" + + return Path( + hf_hub_download( + identity["hf_repo"], + f"{identity['subfolder']}/{filename}", + revision=identity["revision"], + local_dir=str(download_dir), + ) + ) + + _check(fetch("config.json"), identity["config_sha256"]) + index_path = fetch(identity["index_file"]) + _check(index_path, identity["index_sha256"]) + weight_map = json.loads(index_path.read_text())["weight_map"] + + def extract(section: str, filename: str) -> None: + """Pull one manifest section's tensors from their pinned shards, verify, and save.""" + + specs = manifest[section] + tensors = {} + for shard in sorted({weight_map[name] for name in specs}): + if shard not in manifest["weight_shards"]: + raise SystemExit(f"shard {shard} is not pinned in the manifest") + shard_path = fetch(shard) + _check(shard_path, manifest["weight_shards"][shard]["sha256"]) + with safe_open(str(shard_path), framework="pt") as handle: + for name in specs: + if weight_map[name] == shard: + tensors[name] = handle.get_tensor(name).contiguous() + + for name, spec in specs.items(): + tensor = tensors[name] + if str(tensor.dtype).removeprefix("torch.") != spec["dtype"]: + raise SystemExit(f"{name}: dtype {tensor.dtype} != manifest {spec['dtype']}") + if list(tensor.shape) != spec["shape"]: + raise SystemExit(f"{name}: shape {list(tensor.shape)} != manifest {spec['shape']}") + actual = sha256_tensor(tensor) + if actual != spec["sha256"]: + raise SystemExit(f"{name}: sha256 {actual} != manifest {spec['sha256']}") + + target = out_dir / filename + save_file(tensors, str(target), metadata={"revision": identity["revision"]}) + print(f"wrote {len(tensors)} tensors to {target}") + + extract("tensors", EXTRACTED_FILE) + if args.block: + extract("block_tensors", BLOCK_FILE) + print(f"export RL_KERNEL_H3_WEIGHTS={out_dir}") + + +if __name__ == "__main__": + main()