diff --git a/benchmarks/bench_kda_k1_parallelism.py b/benchmarks/bench_kda_k1_parallelism.py new file mode 100644 index 00000000000..29f1e404544 --- /dev/null +++ b/benchmarks/bench_kda_k1_parallelism.py @@ -0,0 +1,488 @@ +# Copyright (c) 2026 by FlashInfer team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Benchmark SM100-family CAKE K1 parallelism against the valid CAKE oracle. + +Every implementation is invoked through ``flashinfer.kda.recurrent_kda``. +Before timing, the script requires M128 owner/helper routes to be +bitwise-identical to CAKE-M128 and checks M64 dual-owner routes against +CAKE-M128 with the recurrent-KDA reference tolerance. Physical routes are then +measured with cold L2. CAKE-M64 participates in the baseline only for its +original fixed ``B=1, H=64`` contract; all other shapes use CAKE-M128. +Optional forced C4/C8 and mailbox-depth configurations preserve the evidence +used to tune dispatch. +""" + +import argparse +import importlib +import json +import random +from contextlib import contextmanager +from pathlib import Path +from typing import Callable, Iterator, Optional + +import numpy as np +import torch + +from flashinfer.kda import recurrent_kda +from flashinfer.kda_prefill import RecurrentKDAPrefillWorkspace +from flashinfer.testing import bench_gpu_time +from flashinfer.utils import get_compute_capability + +kda_prefill = importlib.import_module("flashinfer.kda_prefill") + + +def _parse_forced_config(value: str) -> tuple[int, int]: + try: + cluster_size_text, mailbox_depth_text = value.split(":", 1) + cluster_size = int(cluster_size_text) + mailbox_depth = int(mailbox_depth_text) + except ValueError as error: + raise argparse.ArgumentTypeError( + f"expected CLUSTER_SIZE:MAILBOX_DEPTH, got {value!r}" + ) from error + if cluster_size not in (4, 8): + raise argparse.ArgumentTypeError("forced cluster size must be 4 or 8") + producer_instances = (cluster_size - 1) * 5 + if mailbox_depth <= 0 or mailbox_depth % producer_instances != 0: + raise argparse.ArgumentTypeError( + "forced mailbox depth must be a positive multiple of " + f"{producer_instances} for C{cluster_size}" + ) + return cluster_size, mailbox_depth + + +def _parse_forced_route(value: str) -> tuple[str, int, int]: + try: + variant, cluster_size_text, mailbox_depth_text = value.split(":", 2) + cluster_size = int(cluster_size_text) + mailbox_depth = int(mailbox_depth_text) + except ValueError as error: + raise argparse.ArgumentTypeError( + f"expected VARIANT:CLUSTER_SIZE:MAILBOX_DEPTH, got {value!r}" + ) from error + if variant not in ("m64_k1_parallel", "m128_k1_parallel"): + raise argparse.ArgumentTypeError( + "forced variant must be m64_k1_parallel or m128_k1_parallel" + ) + if cluster_size not in (4, 8): + raise argparse.ArgumentTypeError("forced cluster size must be 4 or 8") + if variant == "m64_k1_parallel" and cluster_size != 4: + raise argparse.ArgumentTypeError( + "m64_k1_parallel has only been validated with cluster size 4" + ) + owner_count = 2 if variant == "m64_k1_parallel" else 1 + producer_instances = (cluster_size - owner_count) * 5 + if mailbox_depth <= 0 or mailbox_depth % producer_instances != 0: + raise argparse.ArgumentTypeError( + "forced mailbox depth must be a positive multiple of " + f"{producer_instances} for {variant} C{cluster_size}" + ) + return variant, cluster_size, mailbox_depth + + +def _parse_varlen_profile(value: str) -> tuple[int, ...]: + try: + lengths = tuple(int(item) for item in value.split(",")) + except ValueError as error: + raise argparse.ArgumentTypeError( + f"expected comma-separated positive lengths, got {value!r}" + ) from error + if not lengths or any(length <= 0 for length in lengths): + raise argparse.ArgumentTypeError( + f"varlen profile lengths must be positive, got {value!r}" + ) + return lengths + + +@contextmanager +def _physical_route( + route: Optional[tuple[str, int, int]], +) -> Iterator[None]: + original = kda_prefill._select_flash_kda_prefill_variant + if route is not None: + kda_prefill._select_flash_kda_prefill_variant = ( + lambda route=route, **_kwargs: route + ) + try: + yield + finally: + kda_prefill._select_flash_kda_prefill_variant = original + + +def _measure( + run: Callable[[], object], + *, + enable_cupti: bool, + warmup_ms: int, + bench_ms: int, +) -> list[float]: + return [ + float(value) + for value in bench_gpu_time( + run, + enable_cupti=enable_cupti, + cold_l2_cache=True, + use_cuda_graph=False, + dry_run_time_ms=warmup_ms, + repeat_time_ms=bench_ms, + ) + ] + + +def _run_case( + *, + batch_size: int, + sequence_length: int, + varlen_profile: Optional[tuple[int, ...]], + num_heads: int, + seed: int, + enable_cupti: bool, + warmup_ms: int, + bench_ms: int, + state_rotations: int, + forced_configs: list[tuple[int, int]], + forced_routes: list[tuple[str, int, int]], + measurement_rounds: int, +) -> dict: + generator = torch.Generator(device="cuda").manual_seed(seed) + packed = varlen_profile is not None + sequence_lengths = ( + varlen_profile + if varlen_profile is not None + else (sequence_length,) * batch_size + ) + num_sequences = len(sequence_lengths) + total_tokens = sum(sequence_lengths) + shape = ( + (1, total_tokens, num_heads, 128) + if packed + else (batch_size, sequence_length, num_heads, 128) + ) + + def randn(dims: tuple[int, ...]) -> torch.Tensor: + return torch.randn(dims, generator=generator, device="cuda").to(torch.bfloat16) + + q, k, v, g = (randn(shape) for _ in range(4)) + beta = randn(shape[:-1]) + A_log = torch.rand((num_heads,), generator=generator, device="cuda") + dt_bias = torch.rand((num_heads, 128), generator=generator, device="cuda") + initial = randn((num_sequences, num_heads, 128, 128)) + offsets = [0] + for length in sequence_lengths: + offsets.append(offsets[-1] + length) + cu_seqlens = ( + torch.tensor(offsets, dtype=torch.int64, device="cuda") if packed else None + ) + seq_order = ( + torch.tensor( + sorted( + range(num_sequences), + key=sequence_lengths.__getitem__, + reverse=True, + ), + dtype=torch.int32, + device="cuda", + ) + if packed + else None + ) + physical_routes = { + f"c{cluster_size}_d{mailbox_depth}": ( + "m128_k1_parallel", + cluster_size, + mailbox_depth, + ) + for cluster_size, mailbox_depth in forced_configs + } + physical_routes.update( + { + f"{variant}_c{cluster_size}_d{mailbox_depth}": ( + variant, + cluster_size, + mailbox_depth, + ) + for variant, cluster_size, mailbox_depth in forced_routes + } + ) + baseline_routes = {"m128": ("m128", 0, 0)} + if not packed and batch_size == 1 and num_heads == 64: + baseline_routes["m64"] = ("m64", 0, 0) + routes = { + **baseline_routes, + **physical_routes, + "k1_parallel": None, + } + outputs = {name: torch.empty_like(q) for name in routes} + states = {name: initial.clone() for name in outputs} + timed_state_pool = ( + initial.unsqueeze(0).expand(state_rotations, *initial.shape).clone() + ) + state_cursor = 0 + workspaces = {name: RecurrentKDAPrefillWorkspace(q.device) for name in outputs} + + def launch(name: str, *, timed: bool = False) -> object: + nonlocal state_cursor + state = states[name] + if timed: + if state_cursor >= state_rotations: + raise RuntimeError( + f"{name} exhausted {state_rotations} preinitialized state slots" + ) + state = timed_state_pool[state_cursor] + state_cursor += 1 + return recurrent_kda( + q=q, + k=k, + v=v, + g=g, + beta=beta, + A_log=A_log, + dt_bias=dt_bias, + scale=128**-0.5, + initial_state=state, + output=outputs[name], + output_final_state=False, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + lower_bound=-5.0, + beta_is_logit=True, + cu_seqlens=cu_seqlens, + seq_order=seq_order, + prefill_workspace=workspaces[name], + ) + + for name, route in routes.items(): + with _physical_route(route): + launch(name) + torch.cuda.synchronize() + helper_routes = [*physical_routes, "k1_parallel"] + auto_route = kda_prefill._select_flash_kda_prefill_variant( + fixed_layout=not packed, + num_sequences=num_sequences, + num_heads=num_heads, + sequence_length=total_tokens if packed else sequence_length, + device=q.device, + ) + for name in helper_routes: + if name not in outputs: + continue + route = routes[name] or auto_route + if route[0] == "m64_k1_parallel": + torch.testing.assert_close( + outputs[name].float(), outputs["m128"].float(), atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + states[name].float(), states["m128"].float(), atol=1e-2, rtol=1e-2 + ) + elif route[0] == "m128_k1_parallel": + torch.testing.assert_close(outputs[name], outputs["m128"], atol=0, rtol=0) + torch.testing.assert_close(states[name], states["m128"], atol=0, rtol=0) + + samples = {name: [] for name in routes} + state_slots_used = {name: [] for name in routes} + route_orders = [] + order_generator = random.Random(seed) + for _ in range(measurement_rounds): + route_order = list(routes) + order_generator.shuffle(route_order) + route_orders.append(route_order) + for name in route_order: + route = routes[name] + if state_cursor: + timed_state_pool[:state_cursor].copy_(initial.unsqueeze(0)) + state_cursor = 0 + torch.cuda.synchronize() + with _physical_route(route): + round_samples = _measure( + lambda name=name: launch(name, timed=True), + enable_cupti=enable_cupti, + warmup_ms=warmup_ms, + bench_ms=bench_ms, + ) + state_slots_used[name].append(state_cursor) + samples[name].extend(round_samples) + + timings = {name: float(np.median(values)) for name, values in samples.items()} + + oracle_ms = min(timings[name] for name in baseline_routes) + forced_results = { + name: { + "cluster_size": route[1], + "mailbox_depth": route[2], + "latency_ms": timings[name], + "speedup_vs_oracle": oracle_ms / timings[name], + } + for name, route in physical_routes.items() + } + return { + "batch_size": 1 if packed else batch_size, + "num_sequences": num_sequences, + "sequence_length": sequence_length if not packed else None, + "sequence_lengths": list(sequence_lengths), + "total_tokens": total_tokens, + "layout": "packed" if packed else "fixed", + "num_heads": num_heads, + "task_count": num_sequences * num_heads, + "m64_ms": timings.get("m64"), + "m128_ms": timings["m128"], + "oracle_ms": oracle_ms, + "k1_parallel_ms": timings["k1_parallel"], + "speedup_vs_oracle": oracle_ms / timings["k1_parallel"], + "forced_results": forced_results, + "samples_ms": samples, + "correctness": ( + "M128 helpers bitwise versus M128; M64 helpers within 1e-2 " + "absolute/relative tolerance versus M128" + ), + "timing_backend": "cupti" if enable_cupti else "cuda_event", + "cold_l2": True, + "cuda_graph": False, + "same_initial_state_per_timed_call": True, + "state_slots_used_per_round": state_slots_used, + "measurement_rounds": measurement_rounds, + "route_orders": route_orders, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--batch-size", type=int, default=1) + parser.add_argument( + "--sequence-lengths", type=int, nargs="+", default=[1024, 2048, 4096, 8192] + ) + parser.add_argument( + "--varlen-profiles", + type=_parse_varlen_profile, + nargs="*", + default=[], + metavar="L0,L1,...", + help=( + "benchmark packed varlen profiles instead of fixed shapes, for " + "example: --varlen-profiles 8192 4096,3072,2048,1024" + ), + ) + parser.add_argument("--num-heads", type=int, nargs="+", default=[4, 8, 16, 24, 32]) + parser.add_argument("--warmup-ms", type=int, default=20) + parser.add_argument("--bench-ms", type=int, default=100) + parser.add_argument("--state-rotations", type=int, default=2048) + parser.add_argument( + "--measurement-rounds", + type=int, + default=1, + help="repeat timings with deterministic shuffled route order", + ) + parser.add_argument( + "--forced-configs", + type=_parse_forced_config, + nargs="*", + default=[], + metavar="C:D", + help=( + "force additional cluster-size/mailbox-depth routes, for " + "example: --forced-configs 4:15 4:30 8:35" + ), + ) + parser.add_argument( + "--forced-routes", + type=_parse_forced_route, + nargs="*", + default=[], + metavar="VARIANT:C:D", + help=( + "force an owner/helper physical variant, for example: " + "--forced-routes m64_k1_parallel:4:10 " + "m128_k1_parallel:8:35" + ), + ) + parser.add_argument("--cupti", action="store_true") + parser.add_argument("--seed", type=int, default=20260813) + parser.add_argument("--json", type=Path) + args = parser.parse_args() + if args.measurement_rounds <= 0: + parser.error("--measurement-rounds must be positive") + + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required") + compute_capability = get_compute_capability(torch.device("cuda")) + if compute_capability not in ((10, 0), (10, 3)): + raise RuntimeError( + "CAKE K1 owner/helper benchmarking requires CC 10.0 or CC 10.3" + ) + properties = torch.cuda.get_device_properties(0) + metadata = { + "device_name": properties.name, + "compute_capability": list(compute_capability), + "multiprocessor_count": properties.multi_processor_count, + "torch_version": torch.__version__, + "torch_cuda_version": torch.version.cuda, + } + print(json.dumps(metadata, sort_keys=True)) + + results = [] + cases = ( + [(sum(profile), profile) for profile in args.varlen_profiles] + if args.varlen_profiles + else [(sequence_length, None) for sequence_length in args.sequence_lengths] + ) + for sequence_length, varlen_profile in cases: + for num_heads in args.num_heads: + result = _run_case( + batch_size=args.batch_size, + sequence_length=sequence_length, + varlen_profile=varlen_profile, + num_heads=num_heads, + seed=args.seed + sequence_length + num_heads, + enable_cupti=args.cupti, + warmup_ms=args.warmup_ms, + bench_ms=args.bench_ms, + state_rotations=args.state_rotations, + forced_configs=args.forced_configs, + forced_routes=args.forced_routes, + measurement_rounds=args.measurement_rounds, + ) + result["hardware"] = metadata + results.append(result) + shape_label = ( + f"B={args.batch_size} T={sequence_length:5d}" + if varlen_profile is None + else "L=" + ",".join(str(length) for length in varlen_profile) + ) + m64_label = ( + f"M64={result['m64_ms']:.6f} ms " + if result["m64_ms"] is not None + else "" + ) + print( + f"{shape_label} H={num_heads:2d} " + f"{m64_label}" + f"M128={result['m128_ms']:.6f} ms " + f"K1={result['k1_parallel_ms']:.6f} ms " + f"speedup={result['speedup_vs_oracle']:.3f}x", + flush=True, + ) + for name, forced in result["forced_results"].items(): + print( + f" {name} depth={forced['mailbox_depth']:2d} " + f"latency={forced['latency_ms']:.6f} ms " + f"speedup={forced['speedup_vs_oracle']:.3f}x", + flush=True, + ) + + if args.json is not None: + args.json.write_text(json.dumps(results, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/csrc/kda/flashkda_bf16_fused_m128_k1_parallel.cu b/csrc/kda/flashkda_bf16_fused_m128_k1_parallel.cu new file mode 100644 index 00000000000..bf8e7e4b3dd --- /dev/null +++ b/csrc/kda/flashkda_bf16_fused_m128_k1_parallel.cu @@ -0,0 +1,2711 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// clang-format off +// Generated by tools/export-generated-programs (device kernel). +// Provenance: generated Loom schedule 'flashkda_bf16_fused_m128'; embedded in the host TU as flashkda_bf16_fused_m128_f0217be48b. +// FlashInfer integration: K1 prep CTAs publish bounded global-mailbox packets to a +// persistent M128 recurrent owner within the same cluster launch. +typedef unsigned char uint8_t; +typedef unsigned short uint16_t; +typedef unsigned int uint32_t; +typedef unsigned long long uint64_t; +typedef signed int int32_t; +typedef short int int16_t; + +#include + +#define LOOM_INF CUDART_INF_F +#define TMEM_NCOLS 256 +#define TMEM_TMEM_STATE_OFFSET 64 +#define TMEM_TMEM_STATE_INP_OFFSET 0 +#define TMEM_TMEM_U_ACC_OFFSET 224 +#define TMEM_TMEM_U2_INP_OFFSET 224 +#define TMEM_TMEM_U2_ACC_OFFSET 0 +#define TMEM_TMEM_OUT_OFFSET 192 +#define TMEM_TMEM_STATE_OUT_OFFSET 64 +#define NUM_CHUNK_PIPE_STAGES 5 +#define SMEM_SMEM_QD_OFF 1024 +#define SMEM_SMEM_QD_STAGE_BYTES 8192 +#define SMEM_SMEM_QD_STRIDE 41984 +#define SMEM_SMEM_G_RAW_OFF 1024 +#define SMEM_SMEM_G_RAW_STAGE_BYTES 8192 +#define SMEM_SMEM_G_RAW_STRIDE 41984 +#define SMEM_SMEM_G_RAW_ALL_OFF 1024 +#define SMEM_SMEM_G_RAW_ALL_STAGE_BYTES 176128 +#define SMEM_SMEM_G_RAW_ALL_STRIDE 176128 +#define SMEM_SMEM_KD_OFF 9216 +#define SMEM_SMEM_KD_STAGE_BYTES 8192 +#define SMEM_SMEM_KD_STRIDE 41984 +#define SMEM_SMEM_Q_RAW_PREFETCH_OFF 17408 +#define SMEM_SMEM_Q_RAW_PREFETCH_STAGE_BYTES 8192 +#define SMEM_SMEM_Q_RAW_PREFETCH_STRIDE 41984 +#define SMEM_SMEM_FINAL_TRANS_OFF 17408 +#define SMEM_SMEM_FINAL_TRANS_STAGE_BYTES 12288 +#define SMEM_SMEM_FINAL_TRANS_STRIDE 41984 +#define SMEM_SMEM_KR_TRANS_OFF 17408 +#define SMEM_SMEM_KR_TRANS_STAGE_BYTES 8192 +#define SMEM_SMEM_KR_TRANS_STRIDE 41984 +#define SMEM_SMEM_MQK_TRANS_OFF 25600 +#define SMEM_SMEM_MQK_TRANS_STAGE_BYTES 2048 +#define SMEM_SMEM_MQK_TRANS_STRIDE 41984 +#define SMEM_SMEM_INV_OFF 29696 +#define SMEM_SMEM_INV_STAGE_BYTES 2048 +#define SMEM_SMEM_INV_STRIDE 41984 +#define SMEM_SMEM_V_OFF 32384 +#define SMEM_SMEM_V_STAGE_BYTES 8192 +#define SMEM_SMEM_V_STRIDE 41984 +#define SMEM_SMEM_KI_OFF 17408 +#define SMEM_SMEM_KI_STAGE_BYTES 8192 +#define SMEM_SMEM_KI_STRIDE 41984 +#define SMEM_SMEM_GATE_OFF 25600 +#define SMEM_SMEM_GATE_STAGE_BYTES 16384 +#define SMEM_SMEM_GATE_STRIDE 41984 +#define SMEM_SMEM_BETA_RAW_OFF 41984 +#define SMEM_SMEM_BETA_RAW_STAGE_BYTES 512 +#define SMEM_SMEM_BETA_RAW_STRIDE 41984 +#define SMEM_SMEM_INV_WORK_OFF 32384 +#define SMEM_SMEM_INV_WORK_STAGE_BYTES 4096 +#define SMEM_SMEM_INV_WORK_STRIDE 41984 +#define SMEM_SMEM_OUT_OFF 210944 +#define SMEM_SMEM_OUT_STAGE_BYTES 8192 +#define SMEM_SMEM_OUT_STRIDE 8192 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_OFF 41984 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_STAGE_BYTES 168452 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_STRIDE 168452 +#define SMEM_SMEM_GT_PREFIX_ALL_OFF 41472 +#define SMEM_SMEM_GT_PREFIX_ALL_STAGE_BYTES 168448 +#define SMEM_SMEM_GT_PREFIX_ALL_STRIDE 168448 +#define SMEM_SMEM_GT_ALL_OFF 31744 +#define SMEM_SMEM_GT_ALL_STAGE_BYTES 168448 +#define SMEM_SMEM_GT_ALL_STRIDE 168448 +#define SMEM_SMEM_PREP_BETA_ALL_OFF 42500 +#define SMEM_SMEM_PREP_BETA_ALL_STAGE_BYTES 168064 +#define SMEM_SMEM_PREP_BETA_ALL_STRIDE 168064 +#define SMEM_SMEM_GATE_RATE_ALL_OFF 42628 +#define SMEM_SMEM_GATE_RATE_ALL_STAGE_BYTES 167940 +#define SMEM_SMEM_GATE_RATE_ALL_STRIDE 167940 +#define SMEM_SMEM_V_ALL_OFF 32384 +#define SMEM_SMEM_V_ALL_STAGE_BYTES 176128 +#define SMEM_SMEM_V_ALL_STRIDE 176128 +#define SMEM_SMEM_GATE_ALL_OFF 25600 +#define SMEM_SMEM_GATE_ALL_STAGE_BYTES 184320 +#define SMEM_SMEM_GATE_ALL_STRIDE 184320 +#define SMEM_TOTAL 227328 +#define THREADS 1024 + +#include + +__device__ __forceinline__ uint32_t elect_sync() { + uint32_t pred = 0; + asm volatile( + "{\n\t" + ".reg .pred %%px;\n\t" + "elect.sync _|%%px, %1;\n\t" + "@%%px mov.s32 %0, 1;\n\t" + "}\n" + : "+r"(pred) + : "r"(0xFFFFFFFF)); + return pred; +} + + +__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) { + asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" + :: "r"(mbar_addr), "r"(count)); +} + + +__device__ __forceinline__ uint32_t mbarrier_try_wait(int mbar_addr, int phase) { + uint32_t token; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64" + " P1, [%1], %2;\n\t" + "selp.u32 %0, 1, 0, P1;\n\t" + "}\n" + : "=r"(token) + : "r"(mbar_addr), "r"(phase) : "memory"); + return token; +} + +__device__ __forceinline__ uint32_t mbarrier_try_wait_cluster(int mbar_addr, int phase) { + uint32_t token; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "mbarrier.try_wait.parity.acquire.cluster.shared::cta.b64" + " P1, [%1], %2;\n\t" + "selp.u32 %0, 1, 0, P1;\n\t" + "}\n" + : "=r"(token) + : "r"(mbar_addr), "r"(phase) : "memory"); + return token; +} + +__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) { + uint32_t ticks = 0x989680; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "LAB_WAIT:\n\t" + "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64" + " P1, [%0], %1, %2;\n\t" + "@P1 bra.uni DONE;\n\t" + "bra.uni LAB_WAIT;\n\t" + "DONE:\n\t" + "}\n" + :: "r"(mbar_addr), "r"(phase), "r"(ticks) : "memory"); +} + +__device__ __forceinline__ void mbarrier_wait_cluster(int mbar_addr, int phase) { + uint32_t ticks = 0x989680; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "LAB_WAIT_CLUSTER:\n\t" + "mbarrier.try_wait.parity.acquire.cluster.shared::cta.b64" + " P1, [%0], %1, %2;\n\t" + "@P1 bra.uni DONE_CLUSTER;\n\t" + "bra.uni LAB_WAIT_CLUSTER;\n\t" + "DONE_CLUSTER:\n\t" + "}\n" + :: "r"(mbar_addr), "r"(phase), "r"(ticks) : "memory"); +} + +__device__ __forceinline__ void mbarrier_wait_token(int mbar_addr, int phase, uint32_t token) { + if (token == 0) { + mbarrier_wait(mbar_addr, phase); + } +} + +__device__ __forceinline__ void mbarrier_wait_token_cluster(int mbar_addr, int phase, uint32_t token) { + if (token == 0) { + mbarrier_wait_cluster(mbar_addr, phase); + } +} + + +__device__ __forceinline__ void tcgen05_mma_f16( + int taddr, uint64_t a_desc, uint64_t b_desc, + uint32_t i_desc, int enable_input_d) { + asm volatile( + "{\n\t" + ".reg .pred p;\n\t" + "setp.ne.b32 p, %4, 0;\n\t" + "tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, p;\n\t" + "}\n" + :: "r"(taddr), "l"(a_desc), "l"(b_desc), + "r"(i_desc), "r"(enable_input_d)); +} + + +__device__ __forceinline__ uint64_t desc_encode(uint64_t x) { + return (x & 0x3FFFFULL) >> 4ULL; +} + + +__device__ __forceinline__ void mma_ts_step( + int taddr_out, int taddr_a, int b_lo, uint32_t b_dhi, + uint32_t i_desc, int enable_d) { + asm volatile( + "{\n\t" + ".reg .pred leader, p;\n\t" + ".reg .b32 dhi;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p, %5, 0;\n\t" + "mov.b32 dhi, %3;\n\t" + "mov.b64 db, {%2, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%1], db, %4, p;\n\t" + "}\n" + :: "r"(taddr_out), "r"(taddr_a), "r"(b_lo), "r"(b_dhi), + "r"(i_desc), "r"(enable_d)); +} + + +__device__ __forceinline__ void elect_commit(int mbar_addr) { + asm volatile( + "{\n\t" + ".reg .pred leader;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];\n\t" + "}\n" + :: "r"(mbar_addr)); +} + + +__device__ __forceinline__ void mbarrier_arrive(int mbar_addr) { + asm volatile( + "mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];" + :: "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void mbarrier_arrive_expect_tx(int mbar_addr, uint32_t bytes) { + asm volatile( + "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" + :: "r"(mbar_addr), "r"(bytes) : "memory"); +} + + +__device__ __forceinline__ void tmem_ld_x32(float* dst, int tmem_addr) { + asm volatile( + "tcgen05.ld.sync.aligned.32x32b.x32.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7," + " %8, %9, %10, %11, %12, %13, %14, %15," + " %16, %17, %18, %19, %20, %21, %22, %23," + " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];" + : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]), + "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7]), + "=f"(dst[8]), "=f"(dst[9]), "=f"(dst[10]), "=f"(dst[11]), + "=f"(dst[12]), "=f"(dst[13]), "=f"(dst[14]), "=f"(dst[15]), + "=f"(dst[16]), "=f"(dst[17]), "=f"(dst[18]), "=f"(dst[19]), + "=f"(dst[20]), "=f"(dst[21]), "=f"(dst[22]), "=f"(dst[23]), + "=f"(dst[24]), "=f"(dst[25]), "=f"(dst[26]), "=f"(dst[27]), + "=f"(dst[28]), "=f"(dst[29]), "=f"(dst[30]), "=f"(dst[31]) + : "r"(tmem_addr)); +} + + +__device__ __forceinline__ void tmem_ld_x16(float* dst, int tmem_addr) { + asm volatile( + "tcgen05.ld.sync.aligned.32x32b.x16.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7," + " %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]), + "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7]), + "=f"(dst[8]), "=f"(dst[9]), "=f"(dst[10]), "=f"(dst[11]), + "=f"(dst[12]), "=f"(dst[13]), "=f"(dst[14]), "=f"(dst[15]) + : "r"(tmem_addr)); +} + + +__device__ __forceinline__ void tmem_st_x32_f32(int tmem_addr, const float* src) { + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x32.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8," + " %9, %10, %11, %12, %13, %14, %15, %16," + " %17, %18, %19, %20, %21, %22, %23, %24," + " %25, %26, %27, %28, %29, %30, %31, %32};" + :: "r"(tmem_addr), + "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]), + "f"(src[4]), "f"(src[5]), "f"(src[6]), "f"(src[7]), + "f"(src[8]), "f"(src[9]), "f"(src[10]), "f"(src[11]), + "f"(src[12]), "f"(src[13]), "f"(src[14]), "f"(src[15]), + "f"(src[16]), "f"(src[17]), "f"(src[18]), "f"(src[19]), + "f"(src[20]), "f"(src[21]), "f"(src[22]), "f"(src[23]), + "f"(src[24]), "f"(src[25]), "f"(src[26]), "f"(src[27]), + "f"(src[28]), "f"(src[29]), "f"(src[30]), "f"(src[31])); +} + + +__device__ __forceinline__ void mbarrier_init_pred(int mbar_addr, uint32_t count, uint32_t pred) { + asm volatile( + "{\n\t" + ".reg .pred p;\n\t" + "setp.ne.b32 p, %2, 0;\n\t" + "@p mbarrier.init.shared::cta.b64 [%0], %1;\n\t" + "}\n" :: "r"(mbar_addr), "r"(count), "r"(pred)); +} + + +__device__ __forceinline__ float approx_exp2(float x) { + float y; + asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x)); + return y; +} + + +__device__ __forceinline__ void fma_f32x2_inplace(float2* a, float2 b, float2 c) { + unsigned long long r; + asm("fma.rn.ftz.f32x2 %0, %1, %2, %3;" + : "=l"(r) + : "l"(*(unsigned long long*)a), "l"(*(unsigned long long*)&b), + "l"(*(unsigned long long*)&c)); + *(unsigned long long*)a = r; +} + +__device__ __forceinline__ void mul_f32x2_inplace(float2* a, float2 b) { + asm("mul.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ void add_f32x2_inplace(float2* a, float2 b) { + asm("add.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ void sub_f32x2_inplace(float2* a, float2 b) { + asm("sub.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ float2 add_f32x2(float2 a, float2 b) { + float2 r; + asm("add.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +__device__ __forceinline__ float2 sub_f32x2(float2 a, float2 b) { + float2 r; + asm("sub.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +__device__ __forceinline__ void fma_scale_x32( + float* sv, const float2* scale2, const float2* neg_max2) +{ + float2* sv_2 = reinterpret_cast(sv); + #pragma unroll + for (int j = 0; j < 16; j++) + fma_f32x2_inplace(&sv_2[j], *scale2, *neg_max2); +} + +__device__ __forceinline__ float2 fma_f32x2(float2 a, float2 b, float2 c) { + float2 r; + asm("fma.rn.ftz.f32x2 %0, %1, %2, %3;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b), + "l"(*(unsigned long long*)&c)); + return r; +} + +__device__ __forceinline__ float2 mul_f32x2(float2 a, float2 b) { + float2 r; + asm("mul.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +// ex2_emulation_f32x2 defined in softmax_frag_exp2_cast helper (or standalone) + + +__device__ __forceinline__ void elect_commit2(int mbar_addr0, int mbar_addr1) { + asm volatile( + "{\n\t" + ".reg .pred leader;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%1];\n\t" + "}\n" + :: "r"(mbar_addr0), "r"(mbar_addr1) : "memory"); +} + + +__device__ __forceinline__ void fence_async_shared() { + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); +} + + +__device__ __forceinline__ uint64_t make_smem_desc(int addr) { + const int SBO = 1024; + return desc_encode(addr) + | (desc_encode(SBO) << 32ULL) + | (1ULL << 46ULL) + | (2ULL << 61ULL); +} + + +__device__ __forceinline__ void tma_3d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.3d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3, %4}], [%5];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_2d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.2d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3}], [%4];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_4d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int z, int w, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.4d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3, %4, %5}], [%6];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(w), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_store_4d( + const void *tmap, int x, int y, int z, int w, unsigned smem_addr) { + asm volatile( + "cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group" + " [%0, {%1, %2, %3, %4}], [%5];" + :: "l"(tmap), "r"(x), "r"(y), "r"(z), "r"(w), "r"(smem_addr) : "memory"); +} + +__device__ __forceinline__ void tcgen05_commit(int mbar_addr) { + asm volatile( + "tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];" + :: "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tmem_st_x8_u32(int addr, const uint32_t* src) { + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x8.b32" + " [%0], {%1,%2,%3,%4,%5,%6,%7,%8};" + :: "r"(addr), + "r"(src[0]), "r"(src[1]), "r"(src[2]), "r"(src[3]), + "r"(src[4]), "r"(src[5]), "r"(src[6]), "r"(src[7])); +} + +__device__ __forceinline__ void tmem_ld_16x256b_x4(float* dst, int addr) { + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x4.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=r"(*reinterpret_cast(&dst[0])), + "=r"(*reinterpret_cast(&dst[1])), + "=r"(*reinterpret_cast(&dst[2])), + "=r"(*reinterpret_cast(&dst[3])), + "=r"(*reinterpret_cast(&dst[4])), + "=r"(*reinterpret_cast(&dst[5])), + "=r"(*reinterpret_cast(&dst[6])), + "=r"(*reinterpret_cast(&dst[7])), + "=r"(*reinterpret_cast(&dst[8])), + "=r"(*reinterpret_cast(&dst[9])), + "=r"(*reinterpret_cast(&dst[10])), + "=r"(*reinterpret_cast(&dst[11])), + "=r"(*reinterpret_cast(&dst[12])), + "=r"(*reinterpret_cast(&dst[13])), + "=r"(*reinterpret_cast(&dst[14])), + "=r"(*reinterpret_cast(&dst[15])) + : "r"(addr) : "memory"); +} + + +__device__ __forceinline__ uint32_t make_warp_uniform(uint32_t val) { + uint32_t result; + asm volatile("shfl.sync.idx.b32 %0, %1, 0, 0x1f, 0xffffffff;" + : "=r"(result) : "r"(val)); + return result; +} + +__device__ __forceinline__ int cluster_rank() { + uint32_t rank; + asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank)); + return static_cast(rank); +} + +__device__ __forceinline__ void cluster_sync() { + asm volatile("barrier.cluster.arrive.aligned;" ::: "memory"); + asm volatile("barrier.cluster.wait.aligned;" ::: "memory"); +} + +constexpr int kK1PacketBytes = 31520; + +__device__ __forceinline__ void publish_k1_to_global( + uint32_t local_stage, unsigned char* packet, unsigned int* flag, + unsigned int ready_value, int prep_tid) { + if (prep_tid != 0) { + return; + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet), "r"(local_stage), "n"(28672) : "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet + 28672), "r"(local_stage + 28672), "n"(2688) + : "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet + 31360), "r"(local_stage + 41472), "n"(160) + : "memory"); + asm volatile("cp.async.bulk.commit_group;" ::: "memory"); + asm volatile("cp.async.bulk.wait_group 0;" ::: "memory"); + asm volatile("st.global.release.gpu.u32 [%0], %1;" + :: "l"(flag), "r"(ready_value) : "memory"); +} + +__device__ __forceinline__ void wait_k1_global_flag( + const unsigned int* flag, unsigned int expected) { + unsigned int ready = 0; + do { + asm volatile("ld.global.acquire.gpu.u32 %0, [%1];" + : "=r"(ready) : "l"(flag) : "memory"); + } while (ready != expected); +} + +__device__ __forceinline__ void store_k1_global_flag( + unsigned int* flag, unsigned int value) { + asm volatile("st.global.release.gpu.u32 [%0], %1;" + :: "l"(flag), "r"(value) : "memory"); +} + +__device__ __forceinline__ void load_k1_from_global( + uint32_t local_stage, uint32_t qk_barrier, + const unsigned char* packet) { + mbarrier_arrive_expect_tx(qk_barrier, kK1PacketBytes); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage), "l"(packet), "n"(28672), "r"(qk_barrier) + : "memory"); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage + 28672), "l"(packet + 28672), "n"(2688), + "r"(qk_barrier) : "memory"); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage + 41472), "l"(packet + 31360), "n"(160), + "r"(qk_barrier) : "memory"); +} + +extern "C" { + +__global__ __launch_bounds__(1024) void +// FLASHINFER INTEGRATION BEGIN: allow exact state alias +kernel_flashkda_bf16_fused_m128(__nv_bfloat16* __restrict__ q, const void* __restrict__ q_tma, __nv_bfloat16* __restrict__ k, const void* __restrict__ k_tma, __nv_bfloat16* __restrict__ v, const void* __restrict__ v_tma, __nv_bfloat16* __restrict__ g, const void* __restrict__ g_tma, __nv_bfloat16* __restrict__ beta, const void* __restrict__ beta_tma, float* __restrict__ A_log, float* __restrict__ dt_bias, long long* __restrict__ cu_seqlens, int* __restrict__ seq_order, __nv_bfloat16* initial_state, __nv_bfloat16* __restrict__ out, const void* __restrict__ out_tma, __nv_bfloat16* final_state, unsigned char* k1_workspace, unsigned int* k1_flags, int mailbox_depth, int cluster_size, int num_heads, int use_initial_state, int store_final_state, float scale, float lower_bound) +// FLASHINFER INTEGRATION END: allow exact state alias +{ + // FLASHINFER INTEGRATION BEGIN: acquire global tensor maps + // CUDA kernel-start ordering does not acquire the tensor-map proxy. + // One thread acquires each 128-byte map; the CTA barrier publishes those + // acquires to every thread before any TMA instruction can use a map. + if (threadIdx.x == 0) { + asm volatile( + "fence.proxy.tensormap::generic.acquire.gpu [%0], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%1], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%2], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%3], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%4], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%5], 128;\n" + :: "l"(q_tma), "l"(k_tma), "l"(v_tma), "l"(g_tma), + "l"(beta_tma), "l"(out_tma) + : "memory"); + } + __syncthreads(); + // FLASHINFER INTEGRATION END: acquire global tensor maps + const int tid = threadIdx.x; + const int warp = make_warp_uniform(tid / 32); + const int lane = tid % 32; + const int kClusterSize = cluster_size; + constexpr int kProducerFirstRank = 1; + const int kProducerCount = kClusterSize - 1; + + const int cta_rank = cluster_rank(); + const int bid = int(blockIdx.x) / kClusterSize; + + extern __shared__ __align__(1024) char smem_raw[]; + int smem; + smem = (int)(unsigned long long)__cvta_generic_to_shared(smem_raw); + + // Kernel setup ops + __nv_bfloat16* smem_qd = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_qd_addr = smem + 1024; + __nv_bfloat16* smem_g_raw = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_g_raw_addr = smem + 1024; + __nv_bfloat16* smem_g_raw_all = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_g_raw_all_addr = smem + 1024; + __nv_bfloat16* smem_kd = reinterpret_cast<__nv_bfloat16*>(smem_raw + 9216); + const int smem_kd_addr = smem + 9216; + __nv_bfloat16* smem_q_raw_prefetch = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_q_raw_prefetch_addr = smem + 17408; + __nv_bfloat16* smem_final_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_final_trans_addr = smem + 17408; + __nv_bfloat16* smem_kr_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_kr_trans_addr = smem + 17408; + __nv_bfloat16* smem_mqk_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 25600); + const int smem_mqk_trans_addr = smem + 25600; + __nv_bfloat16* smem_inv = reinterpret_cast<__nv_bfloat16*>(smem_raw + 29696); + const int smem_inv_addr = smem + 29696; + __nv_bfloat16* smem_v = reinterpret_cast<__nv_bfloat16*>(smem_raw + 32384); + const int smem_v_addr = smem + 32384; + __nv_bfloat16* smem_ki = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_ki_addr = smem + 17408; + float* smem_gate = reinterpret_cast(smem_raw + 25600); + const int smem_gate_addr = smem + 25600; + __nv_bfloat16* smem_beta_raw = reinterpret_cast<__nv_bfloat16*>(smem_raw + 41984); + const int smem_beta_raw_addr = smem + 41984; + __nv_bfloat16* smem_inv_work = reinterpret_cast<__nv_bfloat16*>(smem_raw + 32384); + const int smem_inv_work_addr = smem + 32384; + __nv_bfloat16* smem_out = reinterpret_cast<__nv_bfloat16*>(smem_raw + 210944); + const int smem_out_addr = smem + 210944; + float* smem_restore_factor_all = reinterpret_cast(smem_raw + 41984); + const int smem_restore_factor_all_addr = smem + 41984; + float* smem_gt_prefix_all = reinterpret_cast(smem_raw + 41472); + const int smem_gt_prefix_all_addr = smem + 41472; + float* smem_gt_all = reinterpret_cast(smem_raw + 31744); + const int smem_gt_all_addr = smem + 31744; + float* smem_prep_beta_all = reinterpret_cast(smem_raw + 42500); + const int smem_prep_beta_all_addr = smem + 42500; + float* smem_gate_rate_all = reinterpret_cast(smem_raw + 42628); + const int smem_gate_rate_all_addr = smem + 42628; + __nv_bfloat16* smem_v_all = reinterpret_cast<__nv_bfloat16*>(smem_raw + 32384); + const int smem_v_all_addr = smem + 32384; + float* smem_gate_all = reinterpret_cast(smem_raw + 25600); + const int smem_gate_all_addr = smem + 25600; + + // Mbarrier init (17 groups, 77 barriers) + // Mbarriers at smem_raw[0..616) + + if (warp == 0) { + uint32_t leader = elect_sync(); + // --- pipeline 'chunk_pipe' --- + // qk_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 0, 1, leader); + mbarrier_init_pred(smem + 8, 1, leader); + mbarrier_init_pred(smem + 16, 1, leader); + mbarrier_init_pred(smem + 24, 1, leader); + mbarrier_init_pred(smem + 32, 1, leader); + // gate_raw_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 40, 1, leader); + mbarrier_init_pred(smem + 48, 1, leader); + mbarrier_init_pred(smem + 56, 1, leader); + mbarrier_init_pred(smem + 64, 1, leader); + mbarrier_init_pred(smem + 72, 1, leader); + // qk_raw_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 80, 1, leader); + mbarrier_init_pred(smem + 88, 1, leader); + mbarrier_init_pred(smem + 96, 1, leader); + mbarrier_init_pred(smem + 104, 1, leader); + mbarrier_init_pred(smem + 112, 1, leader); + // v_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 120, 1, leader); + mbarrier_init_pred(smem + 128, 1, leader); + mbarrier_init_pred(smem + 136, 1, leader); + mbarrier_init_pred(smem + 144, 1, leader); + mbarrier_init_pred(smem + 152, 1, leader); + // v_free: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 160, 4, leader); + mbarrier_init_pred(smem + 168, 4, leader); + mbarrier_init_pred(smem + 176, 4, leader); + mbarrier_init_pred(smem + 184, 4, leader); + mbarrier_init_pred(smem + 192, 4, leader); + const int producer_free_arrivals = 1; + // smem_free: 5 barriers + mbarrier_init_pred(smem + 200, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 208, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 216, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 224, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 232, producer_free_arrivals, leader); + // raw_inputs_free: 5 barriers + mbarrier_init_pred(smem + 240, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 248, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 256, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 264, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 272, producer_free_arrivals, leader); + // state_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 280, 4, leader); + mbarrier_init_pred(smem + 288, 4, leader); + mbarrier_init_pred(smem + 296, 4, leader); + mbarrier_init_pred(smem + 304, 4, leader); + mbarrier_init_pred(smem + 312, 4, leader); + // old_out_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 320, 1, leader); + mbarrier_init_pred(smem + 328, 1, leader); + mbarrier_init_pred(smem + 336, 1, leader); + mbarrier_init_pred(smem + 344, 1, leader); + mbarrier_init_pred(smem + 352, 1, leader); + // u_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 360, 4, leader); + mbarrier_init_pred(smem + 368, 4, leader); + mbarrier_init_pred(smem + 376, 4, leader); + mbarrier_init_pred(smem + 384, 4, leader); + mbarrier_init_pred(smem + 392, 4, leader); + // u2_acc_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 400, 1, leader); + mbarrier_init_pred(smem + 408, 1, leader); + mbarrier_init_pred(smem + 416, 1, leader); + mbarrier_init_pred(smem + 424, 1, leader); + mbarrier_init_pred(smem + 432, 1, leader); + // u2_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 440, 4, leader); + mbarrier_init_pred(smem + 448, 4, leader); + mbarrier_init_pred(smem + 456, 4, leader); + mbarrier_init_pred(smem + 464, 4, leader); + mbarrier_init_pred(smem + 472, 4, leader); + // final_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 480, 1, leader); + mbarrier_init_pred(smem + 488, 1, leader); + mbarrier_init_pred(smem + 496, 1, leader); + mbarrier_init_pred(smem + 504, 1, leader); + mbarrier_init_pred(smem + 512, 1, leader); + // out_empty: 1 barriers, init_count=1 + mbarrier_init_pred(smem + 520, 1, leader); + // tmem_dealloc_ready: 1 barriers, init_count=2 + mbarrier_init_pred(smem + 528, 2, leader); + // prep_diag_ready: 5 barriers, init_count=2 + mbarrier_init_pred(smem + 536, 2, leader); + mbarrier_init_pred(smem + 544, 2, leader); + mbarrier_init_pred(smem + 552, 2, leader); + mbarrier_init_pred(smem + 560, 2, leader); + mbarrier_init_pred(smem + 568, 2, leader); + // prep_inv16_ready: 5 barriers, init_count=2 + mbarrier_init_pred(smem + 576, 2, leader); + mbarrier_init_pred(smem + 584, 2, leader); + mbarrier_init_pred(smem + 592, 2, leader); + mbarrier_init_pred(smem + 600, 2, leader); + mbarrier_init_pred(smem + 608, 2, leader); + asm volatile("fence.mbarrier_init.release.cluster;"); + } + + __syncthreads(); + cluster_sync(); + + // TMEM alloc (256 columns, 256 used) + volatile int* tmem_addr_storage = (volatile int*)(smem_raw + 656); + if (cta_rank == 0 && warp == 0) { + int _tmem_hold = smem + 656; + asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(_tmem_hold), "r"(256) : "memory"); + } else if (cta_rank != 0 && tid == 0) { + tmem_addr_storage[0] = 0; + } + + __syncthreads(); + if (cta_rank == 0) { + asm volatile("tcgen05.fence::after_thread_sync;"); + } + + const int mbar_base = smem; + #define qk_full_addr (mbar_base + 0) + #define gate_raw_full_addr (mbar_base + 40) + #define qk_raw_full_addr (mbar_base + 80) + #define v_full_addr (mbar_base + 120) + #define v_free_addr (mbar_base + 160) + #define smem_free_addr (mbar_base + 200) + #define raw_inputs_free_addr (mbar_base + 240) + #define state_inp_ready_addr (mbar_base + 280) + #define old_out_ready_addr (mbar_base + 320) + #define u_inp_ready_addr (mbar_base + 360) + #define u2_acc_ready_addr (mbar_base + 400) + #define u2_inp_ready_addr (mbar_base + 440) + #define final_ready_addr (mbar_base + 480) + #define out_empty_addr (mbar_base + 520) + #define tmem_dealloc_ready_addr (mbar_base + 528) + #define prep_diag_ready_addr (mbar_base + 536) + #define prep_inv16_ready_addr (mbar_base + 576) + const int taddr = tmem_addr_storage[0]; + + // Kernel post-init ops + const int tmem_tmem_state = taddr + 64; + const int tmem_tmem_state_inp = taddr; + const int tmem_tmem_u_acc = taddr + 224; + const int tmem_tmem_u2_inp = taddr + 224; + const int tmem_tmem_u2_acc = taddr; + const int tmem_tmem_out = taddr + 192; + const int tmem_tmem_state_out = taddr + 64; + + // ---- Register redistribution for WGs split across roles ---- + // Dec phase frees registers before any WG attempts inc. + if (cta_rank == 0 && warp >= 8) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + } + + // ---- Role: compute ---- + if (cta_rank == 0 && warp <= 3) { + asm volatile("setmaxnreg.inc.sync.aligned.u32 168;"); + { // compute_main + int task_idx = bid; + int seq_idx = seq_order[task_idx / num_heads]; + int head_idx = task_idx % num_heads; + long long bos = cu_seqlens[seq_idx]; + long long eos = cu_seqlens[seq_idx + 1]; + int seq_len = (int)(eos - bos); + int num_chunks = (seq_len + 32 - 1) / 32; + int warp_in_wg = warp % 4; + const int tmem_row_base = warp_in_wg * 32 << 16; + int state_row = warp_in_wg * 32 + lane; + int warp_id_in_role = (warp - 0); + int compute_local_warp = warp_id_in_role; + long long state_base = (((long long)seq_idx * (long long)num_heads + (long long)head_idx) * 128 + (long long)state_row) * 128; + #pragma unroll + for (int state_col_block = 0; state_col_block < 4; state_col_block++) { + float state_frag[32]; + state_frag[0] = 0.0f; + state_frag[1] = 0.0f; + state_frag[2] = 0.0f; + state_frag[3] = 0.0f; + state_frag[4] = 0.0f; + state_frag[5] = 0.0f; + state_frag[6] = 0.0f; + state_frag[7] = 0.0f; + state_frag[8] = 0.0f; + state_frag[9] = 0.0f; + state_frag[10] = 0.0f; + state_frag[11] = 0.0f; + state_frag[12] = 0.0f; + state_frag[13] = 0.0f; + state_frag[14] = 0.0f; + state_frag[15] = 0.0f; + state_frag[16] = 0.0f; + state_frag[17] = 0.0f; + state_frag[18] = 0.0f; + state_frag[19] = 0.0f; + state_frag[20] = 0.0f; + state_frag[21] = 0.0f; + state_frag[22] = 0.0f; + state_frag[23] = 0.0f; + state_frag[24] = 0.0f; + state_frag[25] = 0.0f; + state_frag[26] = 0.0f; + state_frag[27] = 0.0f; + state_frag[28] = 0.0f; + state_frag[29] = 0.0f; + state_frag[30] = 0.0f; + state_frag[31] = 0.0f; + if (use_initial_state != 0) { + { + const uint4* _vptr_0 = reinterpret_cast(initial_state + state_base + (long long)(state_col_block * 32)); + uint4 _vld_0[2]; + #pragma unroll + for (int _blk = 0; _blk < 2; _blk++) { + _vld_0[_blk] = _vptr_0[_blk]; + uint32_t* _vpairs_0 = reinterpret_cast(&_vld_0[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&state_frag[0 + _blk * 8 + _pair * 2])[0]), "=f"((&state_frag[0 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_0[_pair])); + } + } + } + { + const uint4* _vptr_1 = reinterpret_cast(initial_state + state_base + (long long)(state_col_block * 32) + 16); + uint4 _vld_1[2]; + #pragma unroll + for (int _blk = 0; _blk < 2; _blk++) { + _vld_1[_blk] = _vptr_1[_blk]; + uint32_t* _vpairs_1 = reinterpret_cast(&_vld_1[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&state_frag[16 + _blk * 8 + _pair * 2])[0]), "=f"((&state_frag[16 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_1[_pair])); + } + } + } + } + tmem_st_x32_f32(taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_block * 32), state_frag); + } + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + unsigned int compute_stage = 0; + unsigned int _phase_qk_full = 0; + unsigned int _phase_v_full = 0; + unsigned int _phase_old_out_ready = 0; + unsigned int _phase_u2_acc_ready = 0; + unsigned int _phase_final_ready = 0; + #pragma unroll 1 + for (int chunk_idx = 0; chunk_idx < num_chunks; chunk_idx++) { + mbarrier_wait_cluster(qk_full_addr + (compute_stage) * 8, _phase_qk_full); + #pragma unroll 1 + for (int state_col_block_1 = 0; state_col_block_1 < 4; state_col_block_1++) { + int state_addr = taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_block_1 * 32); + float _tmem_load_0[32]; + tmem_ld_x32(&_tmem_load_0[0], state_addr); + uint32_t _tmem_load_0_bf16[16]; + #pragma unroll + for (int _lp = 0; _lp < 16; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(_tmem_load_0[_lp*2 + 0], _tmem_load_0[_lp*2+1 + 0])); + _tmem_load_0_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x16.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16};" + :: "r"(taddr + (unsigned int)tmem_row_base + (unsigned int)(state_col_block_1 * 16)), "r"(*reinterpret_cast(&_tmem_load_0_bf16[0])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[1])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[2])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[3])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[4])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[5])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[6])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[7])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[8])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[9])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[10])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[11])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[12])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[13])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[14])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[15])) + : "memory"); + float state_scale[16]; + #pragma unroll + for (int state_half = 0; state_half < 2; state_half++) { + #pragma unroll + for (int state_col = 0; state_col < 16; state_col++) { + state_scale[state_col] = smem_gt_all[compute_stage * 10496 + (unsigned int)(state_col_block_1 * 32) + (unsigned int)(state_half * 16) + (unsigned int)state_col]; + } + #pragma unroll + for (int _ls = 0; _ls < 8; _ls++) + mul_f32x2_inplace(&reinterpret_cast((_tmem_load_0 + state_half * 16))[_ls], reinterpret_cast(state_scale)[_ls]); + } + tmem_st_x32_f32(state_addr, _tmem_load_0); + } + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(state_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(v_full_addr + (compute_stage) * 8, _phase_v_full); + mbarrier_wait(old_out_ready_addr + (compute_stage) * 8, _phase_old_out_ready); + float _tmem_load_1[32]; + tmem_ld_x32(&_tmem_load_1[0], taddr + 224 + (unsigned int)tmem_row_base); + #pragma unroll + for (int residual_half = 0; residual_half < 2; residual_half++) { + float residual_v[16]; + float residual_beta[16]; + #pragma unroll + for (int residual_col = 0; residual_col < 16; residual_col++) { + int token_col = residual_half * 16 + residual_col; + __nv_bfloat16 v_value = smem_v_all[compute_stage * 20992 + (unsigned int)(token_col * 128) + (unsigned int)state_row]; + float _cvt_f32_2 = __bfloat162float(v_value); + residual_v[residual_col] = _cvt_f32_2; + residual_beta[residual_col] = smem_prep_beta_all[compute_stage * 10496 + (unsigned int)token_col]; + } + #pragma unroll + for (int _ls = 0; _ls < 8; _ls++) + sub_f32x2_inplace(&reinterpret_cast(residual_v)[_ls], reinterpret_cast((_tmem_load_1 + residual_half * 16))[_ls]); + #pragma unroll + for (int _ls = 0; _ls < 8; _ls++) + mul_f32x2_inplace(&reinterpret_cast(residual_v)[_ls], reinterpret_cast(residual_beta)[_ls]); + uint32_t residual_v_bf16[8]; + #pragma unroll + for (int _lp = 0; _lp < 8; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(residual_v[_lp*2 + 0], residual_v[_lp*2+1 + 0])); + residual_v_bf16[_lp] = *(uint32_t*)&_bf2; + } + tmem_st_x8_u32(taddr + 224 + (unsigned int)tmem_row_base + (unsigned int)(residual_half * 8), (const uint32_t*)residual_v_bf16); + } + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(v_free_addr + (compute_stage) * 8); + mbarrier_arrive(u_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(u2_acc_ready_addr + (compute_stage) * 8, _phase_u2_acc_ready); + float _tmem_load_2[32]; + tmem_ld_x32(&_tmem_load_2[0], taddr + (unsigned int)tmem_row_base); + uint32_t _tmem_load_2_bf16[16]; + #pragma unroll + for (int _lp = 0; _lp < 16; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(_tmem_load_2[_lp*2 + 0], _tmem_load_2[_lp*2+1 + 0])); + _tmem_load_2_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x16.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16};" + :: "r"(taddr + 224 + (unsigned int)tmem_row_base), "r"(*reinterpret_cast(&_tmem_load_2_bf16[0])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[1])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[2])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[3])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[4])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[5])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[6])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[7])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[8])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[9])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[10])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[11])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[12])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[13])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[14])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[15])) + : "memory"); + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(u2_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(final_ready_addr + (compute_stage) * 8, _phase_final_ready); + compute_stage += 1; + if (compute_stage == 5) { compute_stage = 0; _phase_qk_full ^= 1; _phase_v_full ^= 1; _phase_old_out_ready ^= 1; _phase_u2_acc_ready ^= 1; _phase_final_ready ^= 1; } + } + if (store_final_state != 0) { + #pragma unroll + for (int state_col_block_2 = 0; state_col_block_2 < 4; state_col_block_2++) { + float _tmem_load_3[32]; + tmem_ld_x32(&_tmem_load_3[0], taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_block_2 * 32)); + { + __nv_bfloat162 _pk[8]; + _pk[0] = __floats2bfloat162_rn(_tmem_load_3[0 + 0], _tmem_load_3[0 + 1]); + _pk[1] = __floats2bfloat162_rn(_tmem_load_3[0 + 2], _tmem_load_3[0 + 3]); + _pk[2] = __floats2bfloat162_rn(_tmem_load_3[0 + 4], _tmem_load_3[0 + 5]); + _pk[3] = __floats2bfloat162_rn(_tmem_load_3[0 + 6], _tmem_load_3[0 + 7]); + _pk[4] = __floats2bfloat162_rn(_tmem_load_3[0 + 8], _tmem_load_3[0 + 9]); + _pk[5] = __floats2bfloat162_rn(_tmem_load_3[0 + 10], _tmem_load_3[0 + 11]); + _pk[6] = __floats2bfloat162_rn(_tmem_load_3[0 + 12], _tmem_load_3[0 + 13]); + _pk[7] = __floats2bfloat162_rn(_tmem_load_3[0 + 14], _tmem_load_3[0 + 15]); + *reinterpret_cast(&((__nv_bfloat16*)(final_state + (state_base + (long long)(state_col_block_2 * 32))))[0]) = *reinterpret_cast(&_pk[0]); + *reinterpret_cast(&((__nv_bfloat16*)(final_state + (state_base + (long long)(state_col_block_2 * 32))))[8]) = *reinterpret_cast(&_pk[4]); + } + { + __nv_bfloat162 _pk[8]; + _pk[0] = __floats2bfloat162_rn(_tmem_load_3[16 + 0], _tmem_load_3[16 + 1]); + _pk[1] = __floats2bfloat162_rn(_tmem_load_3[16 + 2], _tmem_load_3[16 + 3]); + _pk[2] = __floats2bfloat162_rn(_tmem_load_3[16 + 4], _tmem_load_3[16 + 5]); + _pk[3] = __floats2bfloat162_rn(_tmem_load_3[16 + 6], _tmem_load_3[16 + 7]); + _pk[4] = __floats2bfloat162_rn(_tmem_load_3[16 + 8], _tmem_load_3[16 + 9]); + _pk[5] = __floats2bfloat162_rn(_tmem_load_3[16 + 10], _tmem_load_3[16 + 11]); + _pk[6] = __floats2bfloat162_rn(_tmem_load_3[16 + 12], _tmem_load_3[16 + 13]); + _pk[7] = __floats2bfloat162_rn(_tmem_load_3[16 + 14], _tmem_load_3[16 + 15]); + *reinterpret_cast(&((__nv_bfloat16*)(final_state + (state_base + (long long)(state_col_block_2 * 32) + 16)))[0]) = *reinterpret_cast(&_pk[0]); + *reinterpret_cast(&((__nv_bfloat16*)(final_state + (state_base + (long long)(state_col_block_2 * 32) + 16)))[8]) = *reinterpret_cast(&_pk[4]); + } + } + } + asm volatile("barrier.sync 10, 128;" ::: "memory"); + if (compute_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive(tmem_dealloc_ready_addr); + } + } + } + // ---- Role: epilogue ---- + } else if (cta_rank == 0 && warp >= 4 && warp <= 7) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + { // epilogue_main + int task_idx_1 = bid; + int seq_idx_1 = seq_order[task_idx_1 / num_heads]; + int head_idx_1 = task_idx_1 % num_heads; + long long bos_1 = cu_seqlens[seq_idx_1]; + long long eos_1 = cu_seqlens[seq_idx_1 + 1]; + int seq_len_1 = (int)(eos_1 - bos_1); + int num_chunks_1 = (seq_len_1 + 32 - 1) / 32; + int warp_id_in_role_1 = (warp - 4); + int epilogue_local_warp = warp_id_in_role_1; + int warp_in_wg_1 = warp % 4; + const int tmem_row_base_1 = warp_in_wg_1 * 32 << 16; + int state_row_1 = warp_in_wg_1 * 32 + lane; + unsigned int epilogue_stage = 0; + unsigned int output_stage = 0; + unsigned int _phase_final_ready_1 = 0; + #pragma unroll 1 + for (int chunk_idx_1 = 0; chunk_idx_1 < num_chunks_1; chunk_idx_1++) { + mbarrier_wait(final_ready_addr + (epilogue_stage) * 8, _phase_final_ready_1); + int chunk_is_full = ((seq_len_1 >= (chunk_idx_1 + 1) * 32) ? 1 : 0); + if (chunk_is_full != 0) { + if (epilogue_local_warp == 0) { + if (chunk_idx_1 >= 2) { + asm volatile("cp.async.bulk.wait_group.read 1;"); + } + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + int out_stage_addr = smem_out_addr + output_stage * 8192; + #pragma unroll + for (int dim_half = 0; dim_half < 2; dim_half++) { + float tmem_fragment[16]; + tmem_ld_16x256b_x4( + tmem_fragment, + taddr + 192 + (unsigned int)tmem_row_base_1 + + dim_half * 1048576); + asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory"); + if (dim_half == 1) { + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0 && elect_sync()) { + mbarrier_arrive(out_empty_addr); + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + } + unsigned int out_packed[8]; + #pragma unroll + for (int _lp = 0; _lp < 8; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn( + make_float2(tmem_fragment[_lp * 2], + tmem_fragment[_lp * 2 + 1])); + out_packed[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int token_group = 0; token_group < 2; token_group++) { + int mtx_idx = lane / 8; + int row_addr = lane & 7; + int dim_base = epilogue_local_warp * 32 + dim_half * 16 + (mtx_idx & 1) * 8; + int token_base = token_group * 16 + mtx_idx / 2 * 8; + int token_addr = token_base + row_addr; + int token_pair = token_addr / 2; + int token_parity = token_addr & 1; + int raw_row = token_pair + dim_base / 64 * 16; + int raw_col = (dim_base & 63 ^ (token_pair & 3) << 4 ^ token_parity << 3) + token_parity * 64; + int stsm_offset = (raw_row * 128 + raw_col) * 2; + const int pack_base = token_group * 4; + uint32_t _stmatrix_addr_0 = static_cast((unsigned long long)(out_stage_addr + stsm_offset)); + asm volatile("stmatrix.sync.aligned.m8n8.x4.trans.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_0), "r"(*reinterpret_cast(&out_packed[pack_base])), "r"(*reinterpret_cast(&out_packed[pack_base + 1])), "r"(*reinterpret_cast(&out_packed[pack_base + 2])), "r"(*reinterpret_cast(&out_packed[pack_base + 3])) + : "memory"); + } + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0) { + if (elect_sync()) { + tma_store_4d( + out_tma, 0, + (int)(bos_1 + (long long)(chunk_idx_1 * 32)), + head_idx_1, 0, out_stage_addr); + } + asm volatile("cp.async.bulk.commit_group;"); + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + output_stage = output_stage ^ 1; + } else { + float _tmem_load_6[32]; + tmem_ld_x32(&_tmem_load_6[0], taddr + 192 + (unsigned int)tmem_row_base_1); + asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory"); + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive(out_empty_addr); + } + } + #pragma unroll + for (int token_col_1 = 0; token_col_1 < 32; token_col_1++) { + long long out_token = bos_1 + (long long)(chunk_idx_1 * 32 + token_col_1); + if (out_token < eos_1) { + long long out_idx = (out_token * (long long)num_heads + (long long)head_idx_1) * 128 + (long long)state_row_1; + out[out_idx] = _tmem_load_6[token_col_1]; + } + } + } + if (epilogue_local_warp == 0 && elect_sync()) { + long long packet_idx = + (long long)task_idx_1 * mailbox_depth + + chunk_idx_1 % mailbox_depth; + store_k1_global_flag(k1_flags + packet_idx, 0); + mbarrier_arrive( + raw_inputs_free_addr + epilogue_stage * 8); + mbarrier_arrive(smem_free_addr + epilogue_stage * 8); + } + epilogue_stage += 1; + if (epilogue_stage == 5) { epilogue_stage = 0; _phase_final_ready_1 ^= 1; } + } + if (epilogue_local_warp == 0) { + asm volatile("cp.async.bulk.wait_group 0;"); + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive(tmem_dealloc_ready_addr); + } + } + } + // ---- Role: mma ---- + } else if (cta_rank == 0 && warp == 9) { + { // mma_main + int task_idx_2 = bid; + int seq_idx_2 = seq_order[task_idx_2 / num_heads]; + long long bos_2 = cu_seqlens[seq_idx_2]; + long long eos_2 = cu_seqlens[seq_idx_2 + 1]; + int seq_len_2 = (int)(eos_2 - bos_2); + int num_chunks_2 = (seq_len_2 + 32 - 1) / 32; + unsigned int mma_stage = 0; + unsigned int _phase_qk_full_1 = 0; + unsigned int _phase_state_inp_ready = 0; + unsigned int _phase_out_empty_0 = 1; + unsigned int _phase_u_inp_ready = 0; + unsigned int _phase_u2_inp_ready = 0; + #pragma unroll 1 + for (int _chunk_idx = 0; _chunk_idx < num_chunks_2; _chunk_idx++) { + mbarrier_wait_cluster(qk_full_addr + (mma_stage) * 8, _phase_qk_full_1); + mbarrier_wait(state_inp_ready_addr + (mma_stage) * 8, _phase_state_inp_ready); + mbarrier_wait(out_empty_addr, _phase_out_empty_0); + _phase_out_empty_0 ^= 1; + int _mma_b_addr_0 = smem_qd_addr + mma_stage * 41984; + int _mma_b_lo_0 = make_warp_uniform((_mma_b_addr_0 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 134743184;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 16], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 24], db, id, p1;\n\t" + "add.u32 blo, blo, 250;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 32], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 40], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 48], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 56], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_out), "r"(_mma_b_lo_0), "r"(tmem_tmem_state_inp), "r"(0)); + int _mma_b_addr_1 = smem_kd_addr + mma_stage * 41984; + int _mma_b_lo_1 = make_warp_uniform((_mma_b_addr_1 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 134743184;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 16], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 24], db, id, p1;\n\t" + "add.u32 blo, blo, 250;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 32], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 40], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 48], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 56], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_u_acc), "r"(_mma_b_lo_1), "r"(tmem_tmem_state_inp), "r"(0)); + elect_commit(old_out_ready_addr + mma_stage * 8); + mbarrier_wait(u_inp_ready_addr + (mma_stage) * 8, _phase_u_inp_ready); + int _mma_b_addr_2 = smem_inv_addr + mma_stage * 41984; + int _mma_b_lo_2 = make_warp_uniform((_mma_b_addr_2 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0xC0004010;\n\t" + "mov.b32 id, 134743184;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 64;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_u2_acc), "r"(_mma_b_lo_2), "r"(tmem_tmem_u2_inp), "r"(0)); + elect_commit(u2_acc_ready_addr + (mma_stage) * 8); + mbarrier_wait(u2_inp_ready_addr + (mma_stage) * 8, _phase_u2_inp_ready); + int _mma_b_addr_3 = smem_final_trans_addr + mma_stage * 41984; + int _mma_b_lo_3 = make_warp_uniform(((_mma_b_addr_3 >> 4) & 0x3FFF) | 0x1000000); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 136905872;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 128;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_state_out), "r"(_mma_b_lo_3), "r"(tmem_tmem_u2_inp), "r"(1)); + elect_commit(final_ready_addr + mma_stage * 8); + mma_stage += 1; + if (mma_stage == 5) { mma_stage = 0; _phase_qk_full_1 ^= 1; _phase_state_inp_ready ^= 1; _phase_u_inp_ready ^= 1; _phase_u2_inp_ready ^= 1; } + } + unsigned int _phase_tmem_dealloc_ready_0 = 0; + mbarrier_wait(tmem_dealloc_ready_addr, _phase_tmem_dealloc_ready_0); + _phase_tmem_dealloc_ready_0 ^= 1; + int _tmem_dealloc_addr = *((volatile int*)tmem_addr_storage); + asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(_tmem_dealloc_addr), "r"(256)); + asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"); + } + // ---- Role: load ---- + } else if (cta_rank == 0 && warp == 10) { + { // load_main + int task_idx_3 = bid; + int seq_idx_3 = seq_order[task_idx_3 / num_heads]; + int head_idx_2 = task_idx_3 % num_heads; + long long bos_3 = cu_seqlens[seq_idx_3]; + long long eos_3 = cu_seqlens[seq_idx_3 + 1]; + int seq_len_3 = (int)(eos_3 - bos_3); + int num_chunks_3 = (seq_len_3 + 32 - 1) / 32; + unsigned int load_stage = 0; + unsigned int _phase_v_free = 1; + unsigned int _phase_qk_full_2 = 0; + #pragma unroll 1 + for (int chunk_idx_2 = 0; chunk_idx_2 < num_chunks_3; chunk_idx_2++) { + mbarrier_wait(v_free_addr + (load_stage) * 8, _phase_v_free); + mbarrier_wait_cluster(qk_full_addr + (load_stage) * 8, _phase_qk_full_2); + int chunk_is_full_1 = ((seq_len_3 >= (chunk_idx_2 + 1) * 32) ? 1 : 0); + if (elect_sync()) { + if (chunk_is_full_1 != 0) { + mbarrier_arrive_expect_tx(v_full_addr + (load_stage) * 8, 8192); + tma_3d_gmem2smem(smem_v_addr + load_stage * 41984, v_tma, 0, head_idx_2, (int)(bos_3 + (long long)(chunk_idx_2 * 32)), v_full_addr + (load_stage) * 8); + } + } + if (chunk_is_full_1 == 0) { + #pragma unroll + for (int v_load_iter = 0; v_load_iter < 16; v_load_iter++) { + int v_item = v_load_iter * 32 + lane; + int row = v_item / 16; + int segment = v_item % 16; + long long token = bos_3 + (long long)(chunk_idx_2 * 32 + row); + int token_valid = ((token < eos_3) ? 1 : 0); + long long v_src = (token * (long long)num_heads + (long long)head_idx_2) * 128 + (long long)(segment * 8); + asm volatile("cp.async.cg.shared::cta.global [%0], [%1], 16, %2;" + :: "r"(smem_v_addr + load_stage * 41984 + (unsigned int)((row * 128 + segment * 8) * 2)), "l"(v + v_src), "r"((token_valid != 0) ? 16 : 0)); + } + asm volatile("cp.async.commit_group;"); + asm volatile("cp.async.wait_group 0;"); + } + asm volatile("barrier.sync 8, 32;" ::: "memory"); + if (elect_sync()) { + if (chunk_is_full_1 == 0) { + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + mbarrier_arrive(v_full_addr + (load_stage) * 8); + } + } + load_stage += 1; + if (load_stage == 5) { load_stage = 0; _phase_v_free ^= 1; _phase_qk_full_2 ^= 1; } + } + } + // ---- Role: persistent owner mailbox ingress coordinator ---- + } else if (cta_rank == 0 && warp == 12) { + int task_idx_4 = bid; + int seq_idx_4 = seq_order[task_idx_4 / num_heads]; + int seq_len_4 = int(cu_seqlens[seq_idx_4 + 1] - cu_seqlens[seq_idx_4]); + int num_chunks_4 = (seq_len_4 + 31) / 32; + unsigned int mailbox_stage = 0; + unsigned int raw_phase = 1; + unsigned int smem_phase = 1; + for (int chunk_idx_3 = 0; chunk_idx_3 < num_chunks_4; + ++chunk_idx_3) { + mbarrier_wait( + raw_inputs_free_addr + mailbox_stage * 8, raw_phase); + mbarrier_wait( + smem_free_addr + mailbox_stage * 8, smem_phase); + long long packet_idx = + (long long)task_idx_4 * mailbox_depth + + chunk_idx_3 % mailbox_depth; + wait_k1_global_flag(k1_flags + packet_idx, 1); + if (elect_sync()) { + load_k1_from_global( + smem_qd_addr + mailbox_stage * 41984, + qk_full_addr + mailbox_stage * 8, + k1_workspace + packet_idx * kK1PacketBytes); + } + mailbox_stage += 1; + if (mailbox_stage == 5) { + mailbox_stage = 0; + raw_phase ^= 1; + smem_phase ^= 1; + } + } + // ---- Role: prep ---- + } else if (cta_rank >= kProducerFirstRank && warp >= 12 && warp <= 31) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + { // prep_main + int task_idx_4 = bid; + int seq_idx_4 = seq_order[task_idx_4 / num_heads]; + int head_idx_3 = task_idx_4 % num_heads; + long long bos_4 = cu_seqlens[seq_idx_4]; + long long eos_4 = cu_seqlens[seq_idx_4 + 1]; + int seq_len_4 = (int)(eos_4 - bos_4); + int num_chunks_4 = (seq_len_4 + 32 - 1) / 32; + int instance_id = (warp - 12) / 4; + int prep_instance = instance_id; + int warp_id_in_role_2 = (warp - 12); + int prep_local_warp = warp_id_in_role_2 - prep_instance * 4; + int prep_tid = prep_local_warp * 32 + lane; + const int first_work = + (cta_rank - kProducerFirstRank) * 5 + prep_instance; + const int producer_stride = kProducerCount * 5; + const int total_work = num_chunks_4; + int num_prep_iters = first_work < total_work + ? (total_work - 1 - first_work) / producer_stride + 1 + : 0; + unsigned int prep_stage = (unsigned int)prep_instance; + int gate_rate_stage_f32 = prep_instance * 10496; + if (prep_tid == 0) { + float _expf_0 = __expf(A_log[head_idx_3]); + smem_gate_rate_all[gate_rate_stage_f32] = _expf_0; + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + unsigned int _phase_raw_inputs_free = 1; + unsigned int _phase_gate_raw_full = 0; + unsigned int _phase_smem_free = 1; + unsigned int _phase_qk_raw_full = 0; + unsigned int _phase_prep_diag_ready = 0; + unsigned int _phase_prep_inv16_ready = 0; + #pragma unroll 1 + for (int prep_iter = 0; prep_iter < num_prep_iters; prep_iter++) { + const int work_idx = prep_iter * producer_stride + first_work; + int chunk_idx_3 = work_idx; + int stage_f32 = prep_stage * 10496; + int stage_bf16 = prep_stage * 20992; + int chunk_is_full_2 = ((seq_len_4 >= (chunk_idx_3 + 1) * 32) ? 1 : 0); + float early_beta_value = 0.0f; + float early_gate0 = 0.0f; + if (chunk_is_full_2 != 0) { + mbarrier_wait(raw_inputs_free_addr + (prep_stage) * 8, _phase_raw_inputs_free); + if (prep_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive_expect_tx(gate_raw_full_addr + (prep_stage) * 8, 8704); + tma_3d_gmem2smem(smem_g_raw_addr + prep_stage * 41984, g_tma, 0, head_idx_3, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), gate_raw_full_addr + (prep_stage) * 8); + tma_2d_gmem2smem(smem_beta_raw_addr + prep_stage * 41984, beta_tma, head_idx_3 / 8 * 8, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), gate_raw_full_addr + (prep_stage) * 8); + mbarrier_arrive_expect_tx(qk_raw_full_addr + (prep_stage) * 8, 16384); + tma_4d_gmem2smem(smem_kd_addr + prep_stage * 41984, k_tma, 0, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), head_idx_3, 0, qk_raw_full_addr + (prep_stage) * 8); + } + } + mbarrier_wait(gate_raw_full_addr + (prep_stage) * 8, _phase_gate_raw_full); + if (prep_local_warp == 2 && lane < 32) { + unsigned int beta_raw_pair[1]; + asm volatile("ld.shared.b32 %0, [%1];" : "=r"(*reinterpret_cast(&beta_raw_pair[0])) : "r"(smem_beta_raw_addr + prep_stage * 41984 + (unsigned int)(lane * 16) + (unsigned int)(head_idx_3 % 8 / 2 * 4))); + float beta_raw_pair_fp32[2]; + #pragma unroll + for (int _pair = 0; _pair < 1; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&beta_raw_pair_fp32[_pair * 2])[0]), "=f"((&beta_raw_pair_fp32[_pair * 2])[1]) + : "r"(beta_raw_pair[_pair + 0])); + } + float beta_logit = beta_raw_pair_fp32[0]; + if (head_idx_3 % 2 != 0) { + beta_logit = beta_raw_pair_fp32[1]; + } + float _tanh_approx_0; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_0) : "f"(beta_logit * 0.5f)); + early_beta_value = _tanh_approx_0 * 0.5f + 0.5f; + } + if (prep_tid < 128) { + float early_gate_rate = smem_gate_rate_all[stage_f32]; + float early_gate_bias = dt_bias[head_idx_3 * 128 + prep_tid]; + __nv_bfloat16 early_gate_raw = smem_g_raw_all[stage_bf16 + prep_tid]; + float _cvt_f32_0 = __bfloat162float(early_gate_raw); + float early_gate_arg = early_gate_rate * (_cvt_f32_0 + early_gate_bias); + float _tanh_approx_1; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_1) : "f"(early_gate_arg * 0.5f)); + float early_gate_sigmoid = _tanh_approx_1 * 0.5f + 0.5f; + early_gate0 = lower_bound * 1.4426950408889634f * early_gate_sigmoid; + } + } + mbarrier_wait(smem_free_addr + (prep_stage) * 8, _phase_smem_free); + if (chunk_is_full_2 != 0) { + if (prep_local_warp == 0) { + if (elect_sync()) { + tma_4d_gmem2smem(smem_q_raw_prefetch_addr + prep_stage * 41984, q_tma, 0, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), head_idx_3, 0, qk_raw_full_addr + (prep_stage) * 8); + } + } + } + if (chunk_is_full_2 == 0) { + #pragma unroll + for (int gate_load_pass = 0; gate_load_pass < 4; gate_load_pass++) { + int gate_load_item = gate_load_pass * 128 + prep_tid; + int gate_load_row = gate_load_item / 16; + int gate_load_segment = gate_load_item % 16; + long long gate_load_token = bos_4 + (long long)(chunk_idx_3 * 32 + gate_load_row); + long long gate_load_base = (gate_load_token * (long long)num_heads + (long long)head_idx_3) * 128 + (long long)(gate_load_segment * 8); + asm volatile("cp.async.cg.shared::cta.global [%0], [%1], 16, %2;" + :: "r"(smem_g_raw_addr + prep_stage * 41984 + (unsigned int)(gate_load_item * 16)), "l"(g + gate_load_base), "r"((gate_load_token < eos_4) ? 16 : 0)); + } + } + if (chunk_is_full_2 == 0) { + asm volatile("cp.async.commit_group;"); + asm volatile("cp.async.wait_group 0;"); + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + } + if (prep_local_warp == 2 && lane < 32) { + float beta_value = early_beta_value; + if (chunk_is_full_2 == 0) { + long long beta_token = bos_4 + (long long)(chunk_idx_3 * 32 + lane); + if (beta_token < eos_4) { + float beta_logit_1 = (float)beta[beta_token * (long long)num_heads + (long long)head_idx_3]; + float _tanh_approx_2; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_2) : "f"(beta_logit_1 * 0.5f)); + beta_value = _tanh_approx_2 * 0.5f + 0.5f; + } + } + smem_prep_beta_all[stage_f32 + lane] = beta_value; + } + if (prep_tid < 128) { + int gate_col = prep_tid; + float gate_rate = smem_gate_rate_all[stage_f32]; + float gate_bias = dt_bias[head_idx_3 * 128 + gate_col]; + float prefix_log2 = 0.0f; + for (int gate_row = 0; gate_row < 32; gate_row++) { + long long gate_token = bos_4 + (long long)(chunk_idx_3 * 32 + gate_row); + float gate_log2 = 0.0f; + int gate_needs_compute = 1; + if (gate_row == 0) { + if (chunk_is_full_2 != 0) { + gate_log2 = early_gate0; + gate_needs_compute = 0; + } + } + if (gate_needs_compute != 0) { + if (gate_token < eos_4) { + __nv_bfloat16 gate_raw = smem_g_raw_all[stage_bf16 + gate_row * 128 + gate_col]; + float _cvt_f32_1 = __bfloat162float(gate_raw); + float gate_arg = gate_rate * (_cvt_f32_1 + gate_bias); + float _tanh_approx_3; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_3) : "f"(gate_arg * 0.5f)); + float gate_sigmoid = _tanh_approx_3 * 0.5f + 0.5f; + gate_log2 = lower_bound * 1.4426950408889634f * gate_sigmoid; + } + } + prefix_log2 += gate_log2; + smem_gate_all[stage_f32 + gate_row * 128 + gate_col] = prefix_log2; + } + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + if (chunk_is_full_2 != 0) { + mbarrier_wait(qk_raw_full_addr + (prep_stage) * 8, _phase_qk_raw_full); + } + if (prep_tid < 128) { + float total_log2 = smem_gt_prefix_all[stage_f32 + prep_tid]; + float _exp2_0 = approx_exp2(total_log2 - lower_bound * 1.4426950408889634f * 16.0f); + smem_restore_factor_all[stage_f32 + prep_tid] = _exp2_0; + } + if (prep_tid == 0) { + float _exp2_1 = approx_exp2(lower_bound * 1.4426950408889634f * 16.0f); + smem_restore_factor_all[stage_f32 + 128] = _exp2_1; + } + #pragma unroll 1 + for (int work_pass = 0; work_pass < 4; work_pass++) { + int work_item = work_pass * 128 + prep_tid; + int row_1 = work_item / 16; + int segment_1 = work_item % 16; + long long token_1 = bos_4 + (long long)(chunk_idx_3 * 32 + row_1); + int token_valid_1 = ((token_1 < eos_4) ? 1 : 0); + long long gmem_base = (token_1 * (long long)num_heads + (long long)head_idx_3) * 128 + (long long)(segment_1 * 8); + float q_raw_vec[8]; + float k_raw_vec[8]; + q_raw_vec[0] = 0.0f; + q_raw_vec[1] = 0.0f; + q_raw_vec[2] = 0.0f; + q_raw_vec[3] = 0.0f; + q_raw_vec[4] = 0.0f; + q_raw_vec[5] = 0.0f; + q_raw_vec[6] = 0.0f; + q_raw_vec[7] = 0.0f; + k_raw_vec[0] = 0.0f; + k_raw_vec[1] = 0.0f; + k_raw_vec[2] = 0.0f; + k_raw_vec[3] = 0.0f; + k_raw_vec[4] = 0.0f; + k_raw_vec[5] = 0.0f; + k_raw_vec[6] = 0.0f; + k_raw_vec[7] = 0.0f; + if (chunk_is_full_2 != 0) { + unsigned int packed[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed[0])), "=r"(*reinterpret_cast(&packed[(0) + 1])), "=r"(*reinterpret_cast(&packed[(0) + 2])), "=r"(*reinterpret_cast(&packed[(0) + 3])) + : "r"((smem_q_raw_prefetch_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32[_pair * 2])[0]), "=f"((&packed_fp32[_pair * 2])[1]) + : "r"(packed[_pair + 0])); + } + #pragma unroll + for (int value_idx = 0; value_idx < 8; value_idx++) { + q_raw_vec[value_idx] = packed_fp32[value_idx]; + } + unsigned int packed_0[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0[0])), "=r"(*reinterpret_cast(&packed_0[(0) + 1])), "=r"(*reinterpret_cast(&packed_0[(0) + 2])), "=r"(*reinterpret_cast(&packed_0[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32[_pair * 2])[0]), "=f"((&packed_0_fp32[_pair * 2])[1]) + : "r"(packed_0[_pair + 0])); + } + #pragma unroll + for (int value_idx_1 = 0; value_idx_1 < 8; value_idx_1++) { + k_raw_vec[value_idx_1] = packed_0_fp32[value_idx_1]; + } + } else if (token_valid_1 != 0) { + { + const uint4* _vptr_0 = reinterpret_cast(q + gmem_base); + uint4 _vld_0[1]; + #pragma unroll + for (int _blk = 0; _blk < 1; _blk++) { + _vld_0[_blk] = _vptr_0[_blk]; + uint32_t* _vpairs_0 = reinterpret_cast(&_vld_0[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&q_raw_vec[0 + _blk * 8 + _pair * 2])[0]), "=f"((&q_raw_vec[0 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_0[_pair])); + } + } + } + { + const uint4* _vptr_1 = reinterpret_cast(k + gmem_base); + uint4 _vld_1[1]; + #pragma unroll + for (int _blk = 0; _blk < 1; _blk++) { + _vld_1[_blk] = _vptr_1[_blk]; + uint32_t* _vpairs_1 = reinterpret_cast(&_vld_1[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&k_raw_vec[0 + _blk * 8 + _pair * 2])[0]), "=f"((&k_raw_vec[0 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_1[_pair])); + } + } + } + } + float q_sum = 0.0f; + float k_sum = 0.0f; + for (int elem_in_segment = 0; elem_in_segment < 8; elem_in_segment++) { + float q_raw = q_raw_vec[elem_in_segment]; + float k_raw = k_raw_vec[elem_in_segment]; + float _fma_0 = __fmaf_rn(q_raw, q_raw, q_sum); + q_sum = _fma_0; + float _fma_1 = __fmaf_rn(k_raw, k_raw, k_sum); + k_sum = _fma_1; + } + float _shfl_xor_0 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 8); + q_sum += _shfl_xor_0; + float _shfl_xor_1 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 8); + k_sum += _shfl_xor_1; + float _shfl_xor_2 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 4); + q_sum += _shfl_xor_2; + float _shfl_xor_3 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 4); + k_sum += _shfl_xor_3; + float _shfl_xor_4 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 2); + q_sum += _shfl_xor_4; + float _shfl_xor_5 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 2); + k_sum += _shfl_xor_5; + float _shfl_xor_6 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 1); + q_sum += _shfl_xor_6; + float _shfl_xor_7 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 1); + k_sum += _shfl_xor_7; + float _rsqrt_0 = rsqrtf(q_sum + 1e-06f); + float q_inv = _rsqrt_0; + float _rsqrt_1 = rsqrtf(k_sum + 1e-06f); + float k_inv = _rsqrt_1; + const float2 _scale2_2 = {q_inv, q_inv}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(q_raw_vec)[_ls], _scale2_2); + const float2 _scale2_3 = {k_inv, k_inv}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(k_raw_vec)[_ls], _scale2_3); + float qd_vec[8]; + float kd_vec[8]; + float ki_vec[8]; + for (int elem_in_segment_1 = 0; elem_in_segment_1 < 8; elem_in_segment_1++) { + int col = segment_1 * 8 + elem_in_segment_1; + float prefix = smem_gate_all[stage_f32 + row_1 * 128 + col]; + float common_log2 = lower_bound * 1.4426950408889634f * 16.0f; + float _exp2_2 = approx_exp2(prefix - common_log2); + float decay = _exp2_2; + qd_vec[elem_in_segment_1] = decay; + kd_vec[elem_in_segment_1] = decay; + ki_vec[elem_in_segment_1] = k_raw_vec[elem_in_segment_1] / decay; + } + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(qd_vec)[_ls], reinterpret_cast(q_raw_vec)[_ls]); + const float2 _scale2_4 = {scale, scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(qd_vec)[_ls], _scale2_4); + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(kd_vec)[_ls], reinterpret_cast(k_raw_vec)[_ls]); + unsigned int packed_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(qd_vec[_lp*2 + 0], qd_vec[_lp*2+1 + 0])); + packed_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word = 0; word < 4; word++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word * 4)), "r"(packed_1[word])); + } + unsigned int packed_0_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(kd_vec[_lp*2 + 0], kd_vec[_lp*2+1 + 0])); + packed_0_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_1 = 0; word_1 < 4; word_1++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_1 * 4)), "r"(packed_0_1[word_1])); + } + unsigned int packed_1_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(ki_vec[_lp*2 + 0], ki_vec[_lp*2+1 + 0])); + packed_1_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_2 = 0; word_2 < 4; word_2++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_2 * 4)), "r"(packed_1_1[word_2])); + } + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + int pair_row_base = prep_local_warp / 2 * 16; + int pair_col_base = prep_local_warp % 2 * 16; + unsigned int a_frag[4]; + unsigned int b_frag[4]; + float acc[8]; + if (pair_row_base >= pair_col_base) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[0]), "=f"(acc[1]), "=f"(acc[2]), "=f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[4]), "=f"(acc[(4) + 1]), "=f"(acc[(4) + 2]), "=f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + int row0 = pair_row_base + lane / 4; + int row1 = row0 + 8; + int col0 = pair_col_base + lane % 4 * 2; + float beta0 = smem_prep_beta_all[stage_f32 + row0]; + float beta1 = smem_prep_beta_all[stage_f32 + row1]; + float seed[8]; + seed[0] = 0.0f; + seed[1] = 0.0f; + seed[2] = 0.0f; + seed[3] = 0.0f; + seed[4] = 0.0f; + seed[5] = 0.0f; + seed[6] = 0.0f; + seed[7] = 0.0f; + if (row0 > col0) { + seed[0] = acc[0] * beta0; + } + if (row0 > col0 + 1) { + seed[1] = acc[1] * beta0; + } + if (row1 > col0) { + seed[2] = acc[2] * beta1; + } + if (row1 > col0 + 1) { + seed[3] = acc[3] * beta1; + } + if (row0 > col0 + 8) { + seed[4] = acc[4] * beta0; + } + if (row0 > col0 + 9) { + seed[5] = acc[5] * beta0; + } + if (row1 > col0 + 8) { + seed[6] = acc[6] * beta1; + } + if (row1 > col0 + 9) { + seed[7] = acc[7] * beta1; + } + unsigned int seed_packed[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(seed[_lp*2 + 0], seed[_lp*2+1 + 0])); + seed_packed[_lp] = *(uint32_t*)&_bf2; + } + int seed_lane_row = lane % 16; + int seed_lane_col = lane / 16 * 8; + int byte_off = (pair_row_base + seed_lane_row) * 128 + (pair_col_base + seed_lane_col) * 2; + int swizzled_off = byte_off ^ (byte_off >> 7 & 7) << 4; + int seed_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off; + uint32_t _stmatrix_addr_5 = static_cast((unsigned long long)seed_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_5), "r"(*reinterpret_cast(&seed_packed[0])), "r"(*reinterpret_cast(&seed_packed[1])), "r"(*reinterpret_cast(&seed_packed[2])), "r"(*reinterpret_cast(&seed_packed[3])) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[0]), "=f"(acc[1]), "=f"(acc[2]), "=f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[4]), "=f"(acc[(4) + 1]), "=f"(acc[(4) + 2]), "=f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + } else { + acc[0] = 0.0f; + acc[1] = 0.0f; + acc[2] = 0.0f; + acc[3] = 0.0f; + acc[4] = 0.0f; + acc[5] = 0.0f; + acc[6] = 0.0f; + acc[7] = 0.0f; + } + int row0_1 = pair_row_base + lane / 4; + int row1_1 = row0_1 + 8; + int col0_1 = pair_col_base + lane % 4 * 2; + float mqk[8]; + mqk[0] = 0.0f; + mqk[1] = 0.0f; + mqk[2] = 0.0f; + mqk[3] = 0.0f; + mqk[4] = 0.0f; + mqk[5] = 0.0f; + mqk[6] = 0.0f; + mqk[7] = 0.0f; + if (row0_1 >= col0_1) { + mqk[0] = acc[0]; + } + if (row0_1 >= col0_1 + 1) { + mqk[1] = acc[1]; + } + if (row1_1 >= col0_1) { + mqk[2] = acc[2]; + } + if (row1_1 >= col0_1 + 1) { + mqk[3] = acc[3]; + } + if (row0_1 >= col0_1 + 8) { + mqk[4] = acc[4]; + } + if (row0_1 >= col0_1 + 9) { + mqk[5] = acc[5]; + } + if (row1_1 >= col0_1 + 8) { + mqk[6] = acc[6]; + } + if (row1_1 >= col0_1 + 9) { + mqk[7] = acc[7]; + } + unsigned int mqk_packed[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(mqk[_lp*2 + 0], mqk[_lp*2+1 + 0])); + mqk_packed[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int publish_pair = 0; publish_pair < 2; publish_pair++) { + int publish_row = pair_col_base + publish_pair * 8 + (lane & 7); + int publish_col = 128 + pair_row_base + lane / 8 * 8; + uint32_t _stmatrix_addr_6 = static_cast((unsigned long long)(smem_final_trans_addr + prep_stage * 41984 + (unsigned int)(publish_col / 64 * 4096 + publish_row * 128 + publish_col % 64 * 2 ^ (publish_col / 64 * 4096 + publish_row * 128 + publish_col % 64 * 2 >> 7 & 7) << 4))); + asm volatile("stmatrix.sync.aligned.m8n8.x2.trans.shared.b16 [%0], {%1, %2};\n" + :: "r"(_stmatrix_addr_6), "r"(*reinterpret_cast(&mqk_packed[publish_pair * 2])), "r"(*reinterpret_cast(&mqk_packed[publish_pair * 2 + 1])) + : "memory"); + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + if (prep_tid < 128) { + float total_log2_1 = smem_gt_prefix_all[stage_f32 + prep_tid]; + float _exp2_3 = approx_exp2(total_log2_1); + smem_gt_all[stage_f32 + prep_tid] = _exp2_3; + } + if (prep_local_warp >= 2) { + int stage_f32_0 = prep_stage * 10496; + float restore_scale = smem_restore_factor_all[stage_f32_0 + 128]; + float restore_factor[8]; + int restore_segment = lane & 15; + #pragma unroll + for (int restore_elem = 0; restore_elem < 8; restore_elem++) { + int restore_col = restore_segment * 8 + restore_elem; + restore_factor[restore_elem] = smem_restore_factor_all[stage_f32_0 + restore_col]; + } + #pragma unroll 1 + for (int restore_pass = 0; restore_pass < 6; restore_pass++) { + int restore_row = 8 + (prep_local_warp - 2) * 12 + restore_pass * 2 + (lane >> 4); + float restore_qd_values[8]; + float restore_kd_values[8]; + float restore_ki_values[8]; + unsigned int packed_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_2[0])), "=r"(*reinterpret_cast(&packed_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_2[(0) + 3])) + : "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_1[_pair * 2])[0]), "=f"((&packed_fp32_1[_pair * 2])[1]) + : "r"(packed_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_2 = 0; value_idx_2 < 8; value_idx_2++) { + restore_qd_values[value_idx_2] = packed_fp32_1[value_idx_2]; + } + unsigned int packed_0_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0_2[0])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32_1[_pair * 2])[0]), "=f"((&packed_0_fp32_1[_pair * 2])[1]) + : "r"(packed_0_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_3 = 0; value_idx_3 < 8; value_idx_3++) { + restore_kd_values[value_idx_3] = packed_0_fp32_1[value_idx_3]; + } + unsigned int packed_1_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_1_2[0])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 3])) + : "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_1_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_1_fp32[_pair * 2])[0]), "=f"((&packed_1_fp32[_pair * 2])[1]) + : "r"(packed_1_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_4 = 0; value_idx_4 < 8; value_idx_4++) { + restore_ki_values[value_idx_4] = packed_1_fp32[value_idx_4]; + } + float restore_kr_values[8]; + #pragma unroll + for (int restore_elem_1 = 0; restore_elem_1 < 8; restore_elem_1++) { + restore_kr_values[restore_elem_1] = restore_ki_values[restore_elem_1] * restore_factor[restore_elem_1]; + } + const float2 _scale2_7 = {restore_scale, restore_scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_qd_values)[_ls], _scale2_7); + const float2 _scale2_8 = {restore_scale, restore_scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_kd_values)[_ls], _scale2_8); + unsigned int packed_2_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_qd_values[_lp*2 + 0], restore_qd_values[_lp*2+1 + 0])); + packed_2_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_3 = 0; word_3 < 4; word_3++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_3 * 4)), "r"(packed_2_1[word_3])); + } + unsigned int packed_3[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kd_values[_lp*2 + 0], restore_kd_values[_lp*2+1 + 0])); + packed_3[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_4 = 0; word_4 < 4; word_4++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_4 * 4)), "r"(packed_3[word_4])); + } + unsigned int packed_4[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kr_values[_lp*2 + 0], restore_kr_values[_lp*2+1 + 0])); + packed_4[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_5 = 0; word_5 < 4; word_5++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kr_trans_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_5 * 4)), "r"(packed_4[word_5])); + } + } + } + if (prep_local_warp == 0) { + int inverse_row = lane; + int diag_block = inverse_row / 8; + int lane_in_diag = lane & 7; + float inv_row[8]; + unsigned int packed_5[4]; + int byte_off_1 = inverse_row * 128 + diag_block * 8 * 2; + int swizzled_off_1 = byte_off_1 ^ (byte_off_1 >> 7 & 7) << 4; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_5[0])), "=r"(*reinterpret_cast(&packed_5[(0) + 1])), "=r"(*reinterpret_cast(&packed_5[(0) + 2])), "=r"(*reinterpret_cast(&packed_5[(0) + 3])) + : "r"(smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1)); + float packed_fp32_2[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_2[_pair * 2])[0]), "=f"((&packed_fp32_2[_pair * 2])[1]) + : "r"(packed_5[_pair + 0])); + } + #pragma unroll + for (int value_idx_5 = 0; value_idx_5 < 8; value_idx_5++) { + inv_row[value_idx_5] = packed_fp32_2[value_idx_5]; + } + #pragma unroll + for (int diag_elem = 0; diag_elem < 8; diag_elem++) { + if (lane_in_diag == diag_elem) { + inv_row[diag_elem] = 1.0f; + } + } + int diag_group_base = lane - lane_in_diag; + #pragma unroll + for (int src_row = 0; src_row < 7; src_row++) { + float row_scale = -inv_row[src_row]; + #pragma unroll + for (int prev_col = 0; prev_col < src_row; prev_col++) { + int pivot_lane = diag_group_base + src_row; + float _shfl_0 = __shfl_sync(0xFFFFFFFF, inv_row[prev_col], pivot_lane); + float pivot = _shfl_0; + if (lane_in_diag > src_row) { + float _fma_2 = __fmaf_rn(row_scale, pivot, inv_row[prev_col]); + inv_row[prev_col] = _fma_2; + } + } + if (lane_in_diag > src_row) { + inv_row[src_row] = row_scale; + } + } + unsigned int packed_0_3[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(inv_row[_lp*2 + 0], inv_row[_lp*2+1 + 0])); + packed_0_3[_lp] = *(uint32_t*)&_bf2; + } + int byte_off_1_1 = inverse_row * 128 + diag_block * 8 * 2; + int swizzled_off_2 = byte_off_1_1 ^ (byte_off_1_1 >> 7 & 7) << 4; + #pragma unroll + for (int word_6 = 0; word_6 < 4; word_6++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"(smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_2 + (unsigned int)(word_6 * 4)), "r"(packed_0_3[word_6])); + } + } + if (prep_local_warp < 2) { + if (elect_sync()) { + mbarrier_arrive(prep_diag_ready_addr + (prep_stage) * 8); + } + mbarrier_wait(prep_diag_ready_addr + (prep_stage) * 8, _phase_prep_diag_ready); + } + if (prep_local_warp < 2) { + int lane_row = lane & 7; + int byte_off_2 = (prep_local_warp * 16 + 8 + lane_row) * 128 + (prep_local_warp * 16 + 8) * 2; + int swizzled_off_3 = byte_off_2 ^ (byte_off_2 >> 7 & 7) << 4; + int d_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3; + int byte_off_0 = (prep_local_warp * 16 + 8 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_1_1 = byte_off_0 ^ (byte_off_0 >> 7 & 7) << 4; + int c_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1_1; + int byte_off_2_1 = (prep_local_warp * 16 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_3_1 = byte_off_2_1 ^ (byte_off_2_1 >> 7 & 7) << 4; + int a_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3_1; + unsigned int d_frag[2]; + unsigned int c_frag[1]; + float dc_acc[4]; + unsigned int dc_bf16[2]; + unsigned int inv_a_frag[1]; + float o_acc[4]; + unsigned int o_bf16[2]; + asm volatile("ldmatrix.sync.aligned.m8n8.x1.shared.b16 {%0}, [%1];\n" + : "=r"(d_frag[0]) + : "r"(d_addr) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x1.shared.b16 {%0}, [%1];\n" + : "=r"(d_frag[1]) + : "r"(d_addr) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x1.trans.shared.b16 {%0}, [%1];\n" + : "=r"(c_frag[0]) + : "r"(c_addr) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k8.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5}, {%6}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc_acc[0]), "=f"(dc_acc[1]), "=f"(dc_acc[2]), "=f"(dc_acc[3]) + : "r"(d_frag[0]), "r"(d_frag[1]), "r"(c_frag[0])); + const float2 _scale2_9 = {-1.0f, -1.0f}; + #pragma unroll + for (int _ls = 0; _ls < 2; _ls++) + mul_f32x2_inplace(&reinterpret_cast(dc_acc)[_ls], _scale2_9); + #pragma unroll + for (int _lp = 0; _lp < 2; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(dc_acc[_lp*2 + 0], dc_acc[_lp*2+1 + 0])); + dc_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile("ldmatrix.sync.aligned.m8n8.x1.trans.shared.b16 {%0}, [%1];\n" + : "=r"(inv_a_frag[0]) + : "r"(a_addr) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k8.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5}, {%6}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o_acc[0]), "=f"(o_acc[1]), "=f"(o_acc[2]), "=f"(o_acc[3]) + : "r"(dc_bf16[0]), "r"(dc_bf16[1]), "r"(inv_a_frag[0])); + #pragma unroll + for (int _lp = 0; _lp < 2; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(o_acc[_lp*2 + 0], o_acc[_lp*2+1 + 0])); + o_bf16[_lp] = *(uint32_t*)&_bf2; + } + int byte_off_4 = (prep_local_warp * 16 + 8 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_5 = byte_off_4 ^ (byte_off_4 >> 7 & 7) << 4; + int o_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_5; + uint32_t _stmatrix_addr_10 = static_cast((unsigned long long)o_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x1.shared.b16 [%0], {%1};\n" + :: "r"(_stmatrix_addr_10), "r"(*reinterpret_cast(&o_bf16[0])) + : "memory"); + if (elect_sync()) { + mbarrier_arrive(prep_inv16_ready_addr + (prep_stage) * 8); + } + mbarrier_wait(prep_inv16_ready_addr + (prep_stage) * 8, _phase_prep_inv16_ready); + } + if (prep_local_warp == 0) { + int lane_row_1 = lane % 16; + int lane_col = lane / 16 * 8; + int byte_off_3 = (16 + lane_row_1) * 128 + (16 + lane_col) * 2; + int swizzled_off_4 = byte_off_3 ^ (byte_off_3 >> 7 & 7) << 4; + int d_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_4; + int byte_off_0_1 = (16 + lane_row_1) * 128 + lane_col * 2; + int swizzled_off_1_2 = byte_off_0_1 ^ (byte_off_0_1 >> 7 & 7) << 4; + int c_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1_2; + int byte_off_2_2 = lane_row_1 * 128 + lane_col * 2; + int swizzled_off_3_2 = byte_off_2_2 ^ (byte_off_2_2 >> 7 & 7) << 4; + int a_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3_2; + unsigned int d32_frag[4]; + unsigned int c32_frag[4]; + float dc32_acc[8]; + unsigned int dc32_bf16[4]; + unsigned int a32_frag[4]; + float o32_acc[8]; + unsigned int o32_bf16[4]; + unsigned int zero32_bf16[4]; + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(d32_frag[0]), "=r"(d32_frag[1]), "=r"(d32_frag[2]), "=r"(d32_frag[3]) + : "r"(d_addr_1) + : "memory"); + int d_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)((16 + lane_col) / 16 * 1024 + (16 + lane_row_1) * 32 + (16 + lane_col) % 16 * 2 ^ ((16 + lane_col) / 16 * 1024 + (16 + lane_row_1) * 32 + (16 + lane_col) % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_11 = static_cast((unsigned long long)d_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_11), "r"(*reinterpret_cast(&d32_frag[0])), "r"(*reinterpret_cast(&d32_frag[1])), "r"(*reinterpret_cast(&d32_frag[2])), "r"(*reinterpret_cast(&d32_frag[3])) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(c32_frag[0]), "=r"(c32_frag[1]), "=r"(c32_frag[2]), "=r"(c32_frag[3]) + : "r"(c_addr_1) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc32_acc[0]), "=f"(dc32_acc[1]), "=f"(dc32_acc[2]), "=f"(dc32_acc[3]) + : "r"(d32_frag[0]), "r"(d32_frag[1]), "r"(d32_frag[2]), "r"(d32_frag[3]), "r"(c32_frag[0]), "r"(c32_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc32_acc[4]), "=f"(dc32_acc[(4) + 1]), "=f"(dc32_acc[(4) + 2]), "=f"(dc32_acc[(4) + 3]) + : "r"(d32_frag[0]), "r"(d32_frag[1]), "r"(d32_frag[2]), "r"(d32_frag[3]), "r"(c32_frag[2]), "r"(c32_frag[(2) + 1])); + const float2 _scale2_12 = {-1.0f, -1.0f}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(dc32_acc)[_ls], _scale2_12); + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(dc32_acc[_lp*2 + 0], dc32_acc[_lp*2+1 + 0])); + dc32_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a32_frag[0]), "=r"(a32_frag[1]), "=r"(a32_frag[2]), "=r"(a32_frag[3]) + : "r"(a_addr_1) + : "memory"); + int a_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)(lane_col / 16 * 1024 + lane_row_1 * 32 + lane_col % 16 * 2 ^ (lane_col / 16 * 1024 + lane_row_1 * 32 + lane_col % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_13 = static_cast((unsigned long long)a_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.trans.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_13), "r"(*reinterpret_cast(&a32_frag[0])), "r"(*reinterpret_cast(&a32_frag[1])), "r"(*reinterpret_cast(&a32_frag[2])), "r"(*reinterpret_cast(&a32_frag[3])) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o32_acc[0]), "=f"(o32_acc[1]), "=f"(o32_acc[2]), "=f"(o32_acc[3]) + : "r"(dc32_bf16[0]), "r"(dc32_bf16[1]), "r"(dc32_bf16[2]), "r"(dc32_bf16[3]), "r"(a32_frag[0]), "r"(a32_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o32_acc[4]), "=f"(o32_acc[(4) + 1]), "=f"(o32_acc[(4) + 2]), "=f"(o32_acc[(4) + 3]) + : "r"(dc32_bf16[0]), "r"(dc32_bf16[1]), "r"(dc32_bf16[2]), "r"(dc32_bf16[3]), "r"(a32_frag[2]), "r"(a32_frag[(2) + 1])); + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(o32_acc[_lp*2 + 0], o32_acc[_lp*2+1 + 0])); + o32_bf16[_lp] = *(uint32_t*)&_bf2; + } + int o_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)(lane_col / 16 * 1024 + (16 + lane_row_1) * 32 + lane_col % 16 * 2 ^ (lane_col / 16 * 1024 + (16 + lane_row_1) * 32 + lane_col % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_14 = static_cast((unsigned long long)o_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_14), "r"(*reinterpret_cast(&o32_bf16[0])), "r"(*reinterpret_cast(&o32_bf16[1])), "r"(*reinterpret_cast(&o32_bf16[2])), "r"(*reinterpret_cast(&o32_bf16[3])) + : "memory"); + #pragma unroll + for (int zero_word = 0; zero_word < 4; zero_word++) { + zero32_bf16[zero_word] = 0; + } + int zero_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)((16 + lane_col) / 16 * 1024 + lane_row_1 * 32 + (16 + lane_col) % 16 * 2 ^ ((16 + lane_col) / 16 * 1024 + lane_row_1 * 32 + (16 + lane_col) % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_15 = static_cast((unsigned long long)zero_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_15), "r"(*reinterpret_cast(&zero32_bf16[0])), "r"(*reinterpret_cast(&zero32_bf16[1])), "r"(*reinterpret_cast(&zero32_bf16[2])), "r"(*reinterpret_cast(&zero32_bf16[3])) + : "memory"); + } else if (prep_local_warp == 1) { + int stage_f32_0_1 = prep_stage * 10496; + float restore_scale_1 = smem_restore_factor_all[stage_f32_0_1 + 128]; + float restore_factor_1[8]; + int restore_segment_1 = lane & 15; + #pragma unroll + for (int restore_elem_2 = 0; restore_elem_2 < 8; restore_elem_2++) { + int restore_col_1 = restore_segment_1 * 8 + restore_elem_2; + restore_factor_1[restore_elem_2] = smem_restore_factor_all[stage_f32_0_1 + restore_col_1]; + } + #pragma unroll 1 + for (int restore_pass_1 = 0; restore_pass_1 < 4; restore_pass_1++) { + int restore_row_1 = restore_pass_1 * 2 + (lane >> 4); + float restore_qd_values_1[8]; + float restore_kd_values_1[8]; + float restore_ki_values_1[8]; + unsigned int packed_6[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_6[0])), "=r"(*reinterpret_cast(&packed_6[(0) + 1])), "=r"(*reinterpret_cast(&packed_6[(0) + 2])), "=r"(*reinterpret_cast(&packed_6[(0) + 3])) + : "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32_3[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_3[_pair * 2])[0]), "=f"((&packed_fp32_3[_pair * 2])[1]) + : "r"(packed_6[_pair + 0])); + } + #pragma unroll + for (int value_idx_6 = 0; value_idx_6 < 8; value_idx_6++) { + restore_qd_values_1[value_idx_6] = packed_fp32_3[value_idx_6]; + } + unsigned int packed_0_4[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0_4[0])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 1])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 2])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32_2[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32_2[_pair * 2])[0]), "=f"((&packed_0_fp32_2[_pair * 2])[1]) + : "r"(packed_0_4[_pair + 0])); + } + #pragma unroll + for (int value_idx_7 = 0; value_idx_7 < 8; value_idx_7++) { + restore_kd_values_1[value_idx_7] = packed_0_fp32_2[value_idx_7]; + } + unsigned int packed_1_3[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_1_3[0])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 1])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 2])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 3])) + : "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_1_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_1_fp32_1[_pair * 2])[0]), "=f"((&packed_1_fp32_1[_pair * 2])[1]) + : "r"(packed_1_3[_pair + 0])); + } + #pragma unroll + for (int value_idx_8 = 0; value_idx_8 < 8; value_idx_8++) { + restore_ki_values_1[value_idx_8] = packed_1_fp32_1[value_idx_8]; + } + float restore_kr_values_1[8]; + #pragma unroll + for (int restore_elem_3 = 0; restore_elem_3 < 8; restore_elem_3++) { + restore_kr_values_1[restore_elem_3] = restore_ki_values_1[restore_elem_3] * restore_factor_1[restore_elem_3]; + } + const float2 _scale2_16 = {restore_scale_1, restore_scale_1}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_qd_values_1)[_ls], _scale2_16); + const float2 _scale2_17 = {restore_scale_1, restore_scale_1}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_kd_values_1)[_ls], _scale2_17); + unsigned int packed_2_2[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_qd_values_1[_lp*2 + 0], restore_qd_values_1[_lp*2+1 + 0])); + packed_2_2[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_7 = 0; word_7 < 4; word_7++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_7 * 4)), "r"(packed_2_2[word_7])); + } + unsigned int packed_3_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kd_values_1[_lp*2 + 0], restore_kd_values_1[_lp*2+1 + 0])); + packed_3_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_8 = 0; word_8 < 4; word_8++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_8 * 4)), "r"(packed_3_1[word_8])); + } + unsigned int packed_4_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kr_values_1[_lp*2 + 0], restore_kr_values_1[_lp*2+1 + 0])); + packed_4_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_9 = 0; word_9 < 4; word_9++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kr_trans_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_9 * 4)), "r"(packed_4_1[word_9])); + } + } + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + long long packet_idx = + (long long)task_idx_4 * mailbox_depth + + chunk_idx_3 % mailbox_depth; + if (prep_tid == 0) { + wait_k1_global_flag(k1_flags + packet_idx, 0); + } + publish_k1_to_global( + smem_qd_addr + prep_stage * 41984, + k1_workspace + packet_idx * kK1PacketBytes, + k1_flags + packet_idx, 1, prep_tid); + if (prep_tid == 0) { + mbarrier_arrive( + raw_inputs_free_addr + prep_stage * 8); + mbarrier_arrive(smem_free_addr + prep_stage * 8); + } + for (int _advance = 0; _advance < 5; _advance++) { + prep_stage += 1; + if (prep_stage == 5) { prep_stage = 0; _phase_raw_inputs_free ^= 1; _phase_smem_free ^= 1; _phase_gate_raw_full ^= 1; _phase_qk_raw_full ^= 1; _phase_prep_diag_ready ^= 1; _phase_prep_inv16_ready ^= 1; } + } + } + } + } + + __syncthreads(); + cluster_sync(); + +} + +} // extern "C" + +// clang-format on diff --git a/csrc/kda/flashkda_bf16_fused_m128_k1_parallel_binding.cu b/csrc/kda/flashkda_bf16_fused_m128_k1_parallel_binding.cu new file mode 100644 index 00000000000..12231aa36a6 --- /dev/null +++ b/csrc/kda/flashkda_bf16_fused_m128_k1_parallel_binding.cu @@ -0,0 +1,158 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + */ + +#include "flashkda_binding_common.cuh" + +#define uint8_t flashkda_k1_parallel_uint8_t +#define uint16_t flashkda_k1_parallel_uint16_t +#define uint32_t flashkda_k1_parallel_uint32_t +#define uint64_t flashkda_k1_parallel_uint64_t +#define int32_t flashkda_k1_parallel_int32_t +#define int16_t flashkda_k1_parallel_int16_t +#include "flashkda_bf16_fused_m128_k1_parallel.cu" +#undef uint8_t +#undef uint16_t +#undef uint32_t +#undef uint64_t +#undef int32_t +#undef int16_t + +namespace flashinfer { +namespace flash_kda { + +constexpr int64_t kK1ParallelPacketBytes = 31520; +static_assert(kK1ParallelPacketBytes == kK1PacketBytes); +static_assert(THREADS == 1024); +static_assert(SMEM_TOTAL == 227328); + +void RunM128K1Parallel(TensorView q, TensorView k, TensorView v, TensorView g, TensorView beta, + TensorView beta_tma, TensorView A_log, TensorView dt_bias, + TensorView cu_seqlens, TensorView seq_order, TensorView initial_state, + TensorView out, TensorView final_state, TensorView descriptor_storage, + TensorView k1_workspace, int64_t prepare_descriptors, int64_t num_heads, + int64_t use_initial_state, int64_t store_final_state, int64_t cluster_size, + int64_t mailbox_depth, double scale, double lower_bound, + int64_t cuda_stream) { + TVM_FFI_ICHECK(cuda_stream >= 0) << "cuda_stream must be a non-negative stream handle"; + TVM_FFI_ICHECK(q.device().device_type == kDLCUDA) << "q must be a CUDA tensor"; + const int32_t device_id = q.device().device_id; + ffi::CUDADeviceGuard device_guard(device_id); + CheckFlashKDATarget(device_id); + + const int64_t num_seqs = + CheckCommonInputs(q, k, v, g, beta, beta_tma, A_log, dt_bias, cu_seqlens, seq_order, + initial_state, out, final_state, descriptor_storage, prepare_descriptors, + num_heads, use_initial_state, store_final_state, scale, lower_bound); + CheckCudaTensor(k1_workspace, "k1_workspace", device_id); + CheckDtype(k1_workspace, "k1_workspace", dl_uint8); + for (const auto& named : { + std::pair(&q, "q"), + std::pair(&k, "k"), + std::pair(&v, "v"), + std::pair(&g, "g"), + std::pair(&beta, "beta"), + std::pair(&beta_tma, "beta_tma"), + std::pair(&A_log, "A_log"), + std::pair(&dt_bias, "dt_bias"), + std::pair(&cu_seqlens, "cu_seqlens"), + std::pair(&seq_order, "seq_order"), + std::pair(&out, "out"), + std::pair(&descriptor_storage, "descriptor_storage"), + }) { + CheckNoOverlap(k1_workspace, "k1_workspace", *named.first, named.second); + } + if (use_initial_state != 0) { + CheckNoOverlap(k1_workspace, "k1_workspace", initial_state, "initial_state"); + } + if (store_final_state != 0) { + CheckNoOverlap(k1_workspace, "k1_workspace", final_state, "final_state"); + } + TVM_FFI_ICHECK(cluster_size == 4 || cluster_size == 8) << "cluster_size must be C4 or C8"; + TVM_FFI_ICHECK(mailbox_depth > 0 && mailbox_depth <= std::numeric_limits::max()) + << "mailbox_depth must be in the positive int32 range"; + const int64_t producer_instances = (cluster_size - 1) * 5; + TVM_FFI_ICHECK(mailbox_depth >= producer_instances && mailbox_depth % producer_instances == 0) + << "mailbox_depth must be a positive multiple of the helper producer count " + << producer_instances << " for C" << cluster_size << "; got " << mailbox_depth; + TVM_FFI_ICHECK(SupportsBetaTmaHeadCount(num_heads)) + << "K1-parallel FlashKDA requires H == 1, H == 4, or H >= 8 and divisible by 8"; + + const int64_t num_tasks = num_seqs * num_heads; + const int64_t packet_count = num_tasks * mailbox_depth; + TVM_FFI_ICHECK(packet_count > 0 && + packet_count <= std::numeric_limits::max() / kK1ParallelPacketBytes) + << "K1 mailbox packet count is out of range"; + const int64_t flag_offset = + (packet_count * kK1ParallelPacketBytes + int64_t{255}) & ~int64_t{255}; + TVM_FFI_ICHECK(packet_count <= (std::numeric_limits::max() - flag_offset) / + static_cast(sizeof(uint32_t))) + << "K1 mailbox flag size is out of range"; + const int64_t required_bytes = + flag_offset + packet_count * static_cast(sizeof(uint32_t)); + TVM_FFI_ICHECK(k1_workspace.numel() >= required_bytes) + << "k1_workspace requires " << required_bytes << " bytes, got " << k1_workspace.numel(); + TVM_FFI_ICHECK(reinterpret_cast(k1_workspace.data_ptr()) % 256 == 0) + << "k1_workspace must be 256-byte aligned"; + + constexpr int32_t kSmemBytes = SMEM_TOTAL; + CheckDynamicSmemCapacity(device_id, kSmemBytes); + CheckCuda(cudaFuncSetAttribute(kernel_flashkda_bf16_fused_m128, + cudaFuncAttributeMaxDynamicSharedMemorySize, kSmemBytes), + "cudaFuncSetAttribute(kernel_flashkda_bf16_fused_m128_k1_parallel)"); + + const cudaStream_t stream = reinterpret_cast(static_cast(cuda_stream)); + const TmaPointers tma = EncodeTmaPointers<128>(q, k, v, g, beta_tma, out, descriptor_storage, + prepare_descriptors, stream); + PackBetaForTmaIfNeeded(beta, beta_tma, num_heads, stream); + + auto* workspace_bytes = static_cast(k1_workspace.data_ptr()); + auto* flags = reinterpret_cast(workspace_bytes + flag_offset); + CheckCuda(cudaMemsetAsync(flags, 0, packet_count * sizeof(uint32_t), stream), + "cudaMemsetAsync(K1 mailbox flags)"); + + const int64_t grid_x_i64 = num_tasks * cluster_size; + TVM_FFI_ICHECK(grid_x_i64 > 0 && grid_x_i64 <= std::numeric_limits::max()) + << "K1-parallel FlashKDA grid.x is out of range: " << grid_x_i64; + + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeClusterDimension; + attribute.val.clusterDim = {static_cast(cluster_size), 1u, 1u}; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(static_cast(grid_x_i64), 1, 1); + config.blockDim = dim3(THREADS, 1, 1); + config.dynamicSmemBytes = kSmemBytes; + config.stream = stream; + config.attrs = &attribute; + config.numAttrs = 1; + + CheckCuda( + cudaLaunchKernelEx( + &config, kernel_flashkda_bf16_fused_m128, reinterpret_cast<__nv_bfloat16*>(q.data_ptr()), + tma.q, reinterpret_cast<__nv_bfloat16*>(k.data_ptr()), tma.k, + reinterpret_cast<__nv_bfloat16*>(v.data_ptr()), tma.v, + reinterpret_cast<__nv_bfloat16*>(g.data_ptr()), tma.g, + reinterpret_cast<__nv_bfloat16*>(beta.data_ptr()), tma.beta, + reinterpret_cast(A_log.data_ptr()), reinterpret_cast(dt_bias.data_ptr()), + reinterpret_cast(cu_seqlens.data_ptr()), + reinterpret_cast(seq_order.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(initial_state.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), tma.out, + reinterpret_cast<__nv_bfloat16*>(final_state.data_ptr()), workspace_bytes, flags, + static_cast(mailbox_depth), static_cast(cluster_size), + static_cast(num_heads), static_cast(use_initial_state), + static_cast(store_final_state), static_cast(scale), + static_cast(lower_bound)), + "kernel_flashkda_bf16_fused_m128_k1_parallel launch"); +} + +} // namespace flash_kda +} // namespace flashinfer + +TVM_FFI_DLL_EXPORT_TYPED_FUNC(run, flashinfer::flash_kda::RunM128K1Parallel); diff --git a/csrc/kda/flashkda_bf16_fused_m64_binding.cu b/csrc/kda/flashkda_bf16_fused_m64_binding.cu index ca751411f33..9b47d3c6651 100644 --- a/csrc/kda/flashkda_bf16_fused_m64_binding.cu +++ b/csrc/kda/flashkda_bf16_fused_m64_binding.cu @@ -57,9 +57,8 @@ void RunM64(TensorView q, TensorView k, TensorView v, TensorView g, TensorView b initial_state, out, final_state, descriptor_storage, prepare_descriptors, num_heads, use_initial_state, store_final_state, scale, lower_bound); TVM_FFI_ICHECK(num_seqs == 1 && num_heads == 64) - << "the M64 FlashKDA variant is specialized for fixed N=1, H=64; got " - "N=" - << num_seqs << ", H=" << num_heads; + << "the M64 FlashKDA variant is specialized for fixed N=1, H=64; got N=" << num_seqs + << ", H=" << num_heads; constexpr int32_t kSmemBytes = SMEM_TOTAL; CheckDynamicSmemCapacity(device_id, kSmemBytes); diff --git a/csrc/kda/flashkda_bf16_fused_m64_k1_parallel.cu b/csrc/kda/flashkda_bf16_fused_m64_k1_parallel.cu new file mode 100644 index 00000000000..ff83b52e848 --- /dev/null +++ b/csrc/kda/flashkda_bf16_fused_m64_k1_parallel.cu @@ -0,0 +1,2725 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// clang-format off +// Generated by tools/export-generated-programs (device kernel). +// Provenance: generated Loom schedule 'flashkda_bf16_fused_m64'; embedded in the host TU as flashkda_bf16_fused_m64_f0217be48b. +// FlashInfer integration: K1 prep CTAs publish bounded global-mailbox packets to +// two persistent M64 recurrent owners within the same cluster launch. +typedef unsigned char uint8_t; +typedef unsigned short uint16_t; +typedef unsigned int uint32_t; +typedef unsigned long long uint64_t; +typedef signed int int32_t; +typedef short int int16_t; + +#include + +#define LOOM_INF CUDART_INF_F +#define TMEM_NCOLS 256 +#define TMEM_TMEM_STATE_OFFSET 64 +#define TMEM_TMEM_STATE_INP_OFFSET 0 +#define TMEM_TMEM_U_ACC_OFFSET 224 +#define TMEM_TMEM_U2_INP_OFFSET 224 +#define TMEM_TMEM_U2_ACC_OFFSET 0 +#define TMEM_TMEM_OUT_OFFSET 192 +#define TMEM_TMEM_STATE_OUT_OFFSET 64 +#define NUM_CHUNK_PIPE_STAGES 5 +#define SMEM_SMEM_QD_OFF 1024 +#define SMEM_SMEM_QD_STAGE_BYTES 8192 +#define SMEM_SMEM_QD_STRIDE 41984 +#define SMEM_SMEM_G_RAW_OFF 1024 +#define SMEM_SMEM_G_RAW_STAGE_BYTES 8192 +#define SMEM_SMEM_G_RAW_STRIDE 41984 +#define SMEM_SMEM_G_RAW_ALL_OFF 1024 +#define SMEM_SMEM_G_RAW_ALL_STAGE_BYTES 176128 +#define SMEM_SMEM_G_RAW_ALL_STRIDE 176128 +#define SMEM_SMEM_KD_OFF 9216 +#define SMEM_SMEM_KD_STAGE_BYTES 8192 +#define SMEM_SMEM_KD_STRIDE 41984 +#define SMEM_SMEM_Q_RAW_PREFETCH_OFF 17408 +#define SMEM_SMEM_Q_RAW_PREFETCH_STAGE_BYTES 8192 +#define SMEM_SMEM_Q_RAW_PREFETCH_STRIDE 41984 +#define SMEM_SMEM_FINAL_TRANS_OFF 17408 +#define SMEM_SMEM_FINAL_TRANS_STAGE_BYTES 12288 +#define SMEM_SMEM_FINAL_TRANS_STRIDE 41984 +#define SMEM_SMEM_KR_TRANS_OFF 17408 +#define SMEM_SMEM_KR_TRANS_STAGE_BYTES 8192 +#define SMEM_SMEM_KR_TRANS_STRIDE 41984 +#define SMEM_SMEM_MQK_TRANS_OFF 25600 +#define SMEM_SMEM_MQK_TRANS_STAGE_BYTES 2048 +#define SMEM_SMEM_MQK_TRANS_STRIDE 41984 +#define SMEM_SMEM_INV_OFF 29696 +#define SMEM_SMEM_INV_STAGE_BYTES 2048 +#define SMEM_SMEM_INV_STRIDE 41984 +#define SMEM_SMEM_V_OFF 32384 +#define SMEM_SMEM_V_STAGE_BYTES 4096 +#define SMEM_SMEM_V_STRIDE 41984 +#define SMEM_SMEM_KI_OFF 17408 +#define SMEM_SMEM_KI_STAGE_BYTES 8192 +#define SMEM_SMEM_KI_STRIDE 41984 +#define SMEM_SMEM_GATE_OFF 25600 +#define SMEM_SMEM_GATE_STAGE_BYTES 16384 +#define SMEM_SMEM_GATE_STRIDE 41984 +#define SMEM_SMEM_BETA_RAW_OFF 41984 +#define SMEM_SMEM_BETA_RAW_STAGE_BYTES 512 +#define SMEM_SMEM_BETA_RAW_STRIDE 41984 +#define SMEM_SMEM_INV_WORK_OFF 32384 +#define SMEM_SMEM_INV_WORK_STAGE_BYTES 4096 +#define SMEM_SMEM_INV_WORK_STRIDE 41984 +#define SMEM_SMEM_OUT_OFF 210944 +#define SMEM_SMEM_OUT_STAGE_BYTES 4096 +#define SMEM_SMEM_OUT_STRIDE 4096 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_OFF 41984 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_STAGE_BYTES 168452 +#define SMEM_SMEM_RESTORE_FACTOR_ALL_STRIDE 168452 +#define SMEM_SMEM_GT_PREFIX_ALL_OFF 41472 +#define SMEM_SMEM_GT_PREFIX_ALL_STAGE_BYTES 168448 +#define SMEM_SMEM_GT_PREFIX_ALL_STRIDE 168448 +#define SMEM_SMEM_GT_ALL_OFF 31744 +#define SMEM_SMEM_GT_ALL_STAGE_BYTES 168448 +#define SMEM_SMEM_GT_ALL_STRIDE 168448 +#define SMEM_SMEM_PREP_BETA_ALL_OFF 42500 +#define SMEM_SMEM_PREP_BETA_ALL_STAGE_BYTES 168064 +#define SMEM_SMEM_PREP_BETA_ALL_STRIDE 168064 +#define SMEM_SMEM_GATE_RATE_ALL_OFF 42628 +#define SMEM_SMEM_GATE_RATE_ALL_STAGE_BYTES 167940 +#define SMEM_SMEM_GATE_RATE_ALL_STRIDE 167940 +#define SMEM_SMEM_GATE_ALL_OFF 25600 +#define SMEM_SMEM_GATE_ALL_STAGE_BYTES 184320 +#define SMEM_SMEM_GATE_ALL_STRIDE 184320 +#define SMEM_TOTAL 219136 +#define THREADS 1024 + +#include + +__device__ __forceinline__ uint32_t elect_sync() { + uint32_t pred = 0; + asm volatile( + "{\n\t" + ".reg .pred %%px;\n\t" + "elect.sync _|%%px, %1;\n\t" + "@%%px mov.s32 %0, 1;\n\t" + "}\n" + : "+r"(pred) + : "r"(0xFFFFFFFF)); + return pred; +} + + +__device__ __forceinline__ void mbarrier_init(int mbar_addr, int count) { + asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" + :: "r"(mbar_addr), "r"(count)); +} + + +__device__ __forceinline__ uint32_t mbarrier_try_wait(int mbar_addr, int phase) { + uint32_t token; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64" + " P1, [%1], %2;\n\t" + "selp.u32 %0, 1, 0, P1;\n\t" + "}\n" + : "=r"(token) + : "r"(mbar_addr), "r"(phase) : "memory"); + return token; +} + +__device__ __forceinline__ uint32_t mbarrier_try_wait_cluster(int mbar_addr, int phase) { + uint32_t token; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "mbarrier.try_wait.parity.acquire.cluster.shared::cta.b64" + " P1, [%1], %2;\n\t" + "selp.u32 %0, 1, 0, P1;\n\t" + "}\n" + : "=r"(token) + : "r"(mbar_addr), "r"(phase) : "memory"); + return token; +} + +__device__ __forceinline__ void mbarrier_wait(int mbar_addr, int phase) { + uint32_t ticks = 0x989680; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "LAB_WAIT:\n\t" + "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64" + " P1, [%0], %1, %2;\n\t" + "@P1 bra.uni DONE;\n\t" + "bra.uni LAB_WAIT;\n\t" + "DONE:\n\t" + "}\n" + :: "r"(mbar_addr), "r"(phase), "r"(ticks) : "memory"); +} + +__device__ __forceinline__ void mbarrier_wait_cluster(int mbar_addr, int phase) { + uint32_t ticks = 0x989680; + asm volatile( + "{\n\t" + ".reg .pred P1;\n\t" + "LAB_WAIT_CLUSTER:\n\t" + "mbarrier.try_wait.parity.acquire.cluster.shared::cta.b64" + " P1, [%0], %1, %2;\n\t" + "@P1 bra.uni DONE_CLUSTER;\n\t" + "bra.uni LAB_WAIT_CLUSTER;\n\t" + "DONE_CLUSTER:\n\t" + "}\n" + :: "r"(mbar_addr), "r"(phase), "r"(ticks) : "memory"); +} + +__device__ __forceinline__ void mbarrier_wait_token(int mbar_addr, int phase, uint32_t token) { + if (token == 0) { + mbarrier_wait(mbar_addr, phase); + } +} + +__device__ __forceinline__ void mbarrier_wait_token_cluster(int mbar_addr, int phase, uint32_t token) { + if (token == 0) { + mbarrier_wait_cluster(mbar_addr, phase); + } +} + + +__device__ __forceinline__ void tcgen05_mma_f16( + int taddr, uint64_t a_desc, uint64_t b_desc, + uint32_t i_desc, int enable_input_d) { + asm volatile( + "{\n\t" + ".reg .pred p;\n\t" + "setp.ne.b32 p, %4, 0;\n\t" + "tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, p;\n\t" + "}\n" + :: "r"(taddr), "l"(a_desc), "l"(b_desc), + "r"(i_desc), "r"(enable_input_d)); +} + + +__device__ __forceinline__ uint64_t desc_encode(uint64_t x) { + return (x & 0x3FFFFULL) >> 4ULL; +} + + +__device__ __forceinline__ void mma_ts_step( + int taddr_out, int taddr_a, int b_lo, uint32_t b_dhi, + uint32_t i_desc, int enable_d) { + asm volatile( + "{\n\t" + ".reg .pred leader, p;\n\t" + ".reg .b32 dhi;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p, %5, 0;\n\t" + "mov.b32 dhi, %3;\n\t" + "mov.b64 db, {%2, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%1], db, %4, p;\n\t" + "}\n" + :: "r"(taddr_out), "r"(taddr_a), "r"(b_lo), "r"(b_dhi), + "r"(i_desc), "r"(enable_d)); +} + + +__device__ __forceinline__ void elect_commit(int mbar_addr) { + asm volatile( + "{\n\t" + ".reg .pred leader;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];\n\t" + "}\n" + :: "r"(mbar_addr)); +} + + +__device__ __forceinline__ void mbarrier_arrive(int mbar_addr) { + asm volatile( + "mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];" + :: "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void mbarrier_arrive_expect_tx(int mbar_addr, uint32_t bytes) { + asm volatile( + "mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;" + :: "r"(mbar_addr), "r"(bytes) : "memory"); +} + + +__device__ __forceinline__ void tmem_ld_x32(float* dst, int tmem_addr) { + asm volatile( + "tcgen05.ld.sync.aligned.32x32b.x32.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7," + " %8, %9, %10, %11, %12, %13, %14, %15," + " %16, %17, %18, %19, %20, %21, %22, %23," + " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];" + : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]), + "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7]), + "=f"(dst[8]), "=f"(dst[9]), "=f"(dst[10]), "=f"(dst[11]), + "=f"(dst[12]), "=f"(dst[13]), "=f"(dst[14]), "=f"(dst[15]), + "=f"(dst[16]), "=f"(dst[17]), "=f"(dst[18]), "=f"(dst[19]), + "=f"(dst[20]), "=f"(dst[21]), "=f"(dst[22]), "=f"(dst[23]), + "=f"(dst[24]), "=f"(dst[25]), "=f"(dst[26]), "=f"(dst[27]), + "=f"(dst[28]), "=f"(dst[29]), "=f"(dst[30]), "=f"(dst[31]) + : "r"(tmem_addr)); +} + + +__device__ __forceinline__ void tmem_ld_x16(float* dst, int tmem_addr) { + asm volatile( + "tcgen05.ld.sync.aligned.32x32b.x16.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7," + " %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=f"(dst[0]), "=f"(dst[1]), "=f"(dst[2]), "=f"(dst[3]), + "=f"(dst[4]), "=f"(dst[5]), "=f"(dst[6]), "=f"(dst[7]), + "=f"(dst[8]), "=f"(dst[9]), "=f"(dst[10]), "=f"(dst[11]), + "=f"(dst[12]), "=f"(dst[13]), "=f"(dst[14]), "=f"(dst[15]) + : "r"(tmem_addr)); +} + + +__device__ __forceinline__ void tmem_st_x32_f32(int tmem_addr, const float* src) { + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x32.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8," + " %9, %10, %11, %12, %13, %14, %15, %16," + " %17, %18, %19, %20, %21, %22, %23, %24," + " %25, %26, %27, %28, %29, %30, %31, %32};" + :: "r"(tmem_addr), + "f"(src[0]), "f"(src[1]), "f"(src[2]), "f"(src[3]), + "f"(src[4]), "f"(src[5]), "f"(src[6]), "f"(src[7]), + "f"(src[8]), "f"(src[9]), "f"(src[10]), "f"(src[11]), + "f"(src[12]), "f"(src[13]), "f"(src[14]), "f"(src[15]), + "f"(src[16]), "f"(src[17]), "f"(src[18]), "f"(src[19]), + "f"(src[20]), "f"(src[21]), "f"(src[22]), "f"(src[23]), + "f"(src[24]), "f"(src[25]), "f"(src[26]), "f"(src[27]), + "f"(src[28]), "f"(src[29]), "f"(src[30]), "f"(src[31])); +} + + +__device__ __forceinline__ void mbarrier_init_pred(int mbar_addr, uint32_t count, uint32_t pred) { + asm volatile( + "{\n\t" + ".reg .pred p;\n\t" + "setp.ne.b32 p, %2, 0;\n\t" + "@p mbarrier.init.shared::cta.b64 [%0], %1;\n\t" + "}\n" :: "r"(mbar_addr), "r"(count), "r"(pred)); +} + + +__device__ __forceinline__ float approx_exp2(float x) { + float y; + asm("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x)); + return y; +} + + +__device__ __forceinline__ void fma_f32x2_inplace(float2* a, float2 b, float2 c) { + unsigned long long r; + asm("fma.rn.ftz.f32x2 %0, %1, %2, %3;" + : "=l"(r) + : "l"(*(unsigned long long*)a), "l"(*(unsigned long long*)&b), + "l"(*(unsigned long long*)&c)); + *(unsigned long long*)a = r; +} + +__device__ __forceinline__ void mul_f32x2_inplace(float2* a, float2 b) { + asm("mul.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ void add_f32x2_inplace(float2* a, float2 b) { + asm("add.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ void sub_f32x2_inplace(float2* a, float2 b) { + asm("sub.rn.ftz.f32x2 %0, %0, %1;" + : "+l"(*(unsigned long long*)a) : "l"(*(unsigned long long*)&b)); +} + +__device__ __forceinline__ float2 add_f32x2(float2 a, float2 b) { + float2 r; + asm("add.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +__device__ __forceinline__ float2 sub_f32x2(float2 a, float2 b) { + float2 r; + asm("sub.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +__device__ __forceinline__ void fma_scale_x32( + float* sv, const float2* scale2, const float2* neg_max2) +{ + float2* sv_2 = reinterpret_cast(sv); + #pragma unroll + for (int j = 0; j < 16; j++) + fma_f32x2_inplace(&sv_2[j], *scale2, *neg_max2); +} + +__device__ __forceinline__ float2 fma_f32x2(float2 a, float2 b, float2 c) { + float2 r; + asm("fma.rn.ftz.f32x2 %0, %1, %2, %3;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b), + "l"(*(unsigned long long*)&c)); + return r; +} + +__device__ __forceinline__ float2 mul_f32x2(float2 a, float2 b) { + float2 r; + asm("mul.rn.ftz.f32x2 %0, %1, %2;" + : "=l"(*(unsigned long long*)&r) + : "l"(*(unsigned long long*)&a), "l"(*(unsigned long long*)&b)); + return r; +} + +// ex2_emulation_f32x2 defined in softmax_frag_exp2_cast helper (or standalone) + + +__device__ __forceinline__ void elect_commit2(int mbar_addr0, int mbar_addr1) { + asm volatile( + "{\n\t" + ".reg .pred leader;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];\n\t" + "@leader tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%1];\n\t" + "}\n" + :: "r"(mbar_addr0), "r"(mbar_addr1) : "memory"); +} + + +__device__ __forceinline__ void fence_async_shared() { + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); +} + + +__device__ __forceinline__ uint64_t make_smem_desc(int addr) { + const int SBO = 1024; + return desc_encode(addr) + | (desc_encode(SBO) << 32ULL) + | (1ULL << 46ULL) + | (2ULL << 61ULL); +} + + +__device__ __forceinline__ void tma_3d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.3d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3, %4}], [%5];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_2d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.2d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3}], [%4];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_4d_gmem2smem( + int dst, const void *tmap_ptr, int x, int y, int z, int w, int mbar_addr) { + asm volatile( + "cp.async.bulk.tensor.4d.shared::cta.global" + ".mbarrier::complete_tx::bytes" + " [%0], [%1, {%2, %3, %4, %5}], [%6];" + :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(w), + "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tma_store_4d( + const void *tmap, int x, int y, int z, int w, unsigned smem_addr) { + asm volatile( + "cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group" + " [%0, {%1, %2, %3, %4}], [%5];" + :: "l"(tmap), "r"(x), "r"(y), "r"(z), "r"(w), "r"(smem_addr) : "memory"); +} + +__device__ __forceinline__ void tcgen05_commit(int mbar_addr) { + asm volatile( + "tcgen05.commit.cta_group::1.mbarrier::arrive::one" + ".shared::cluster.b64 [%0];" + :: "r"(mbar_addr) : "memory"); +} + + +__device__ __forceinline__ void tmem_st_x8_u32(int addr, const uint32_t* src) { + asm volatile( + "tcgen05.st.sync.aligned.32x32b.x8.b32" + " [%0], {%1,%2,%3,%4,%5,%6,%7,%8};" + :: "r"(addr), + "r"(src[0]), "r"(src[1]), "r"(src[2]), "r"(src[3]), + "r"(src[4]), "r"(src[5]), "r"(src[6]), "r"(src[7])); +} + +__device__ __forceinline__ uint32_t make_warp_uniform(uint32_t val) { + uint32_t result; + asm volatile("shfl.sync.idx.b32 %0, %1, 0, 0x1f, 0xffffffff;" + : "=r"(result) : "r"(val)); + return result; +} + +__device__ __forceinline__ int cluster_rank() { + uint32_t rank; + asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank)); + return static_cast(rank); +} + +__device__ __forceinline__ void cluster_sync() { + asm volatile("barrier.cluster.arrive.aligned;" ::: "memory"); + asm volatile("barrier.cluster.wait.aligned;" ::: "memory"); +} + +constexpr int kK1PacketBytes = 31520; + +__device__ __forceinline__ void publish_k1_to_global( + uint32_t local_stage, unsigned char* packet, unsigned int* flag, + unsigned int ready_value, int prep_tid) { + if (prep_tid != 0) { + return; + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet), "r"(local_stage), "n"(28672) : "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet + 28672), "r"(local_stage + 28672), "n"(2688) + : "memory"); + asm volatile( + "cp.async.bulk.global.shared::cta.bulk_group [%0], [%1], %2;" + :: "l"(packet + 31360), "r"(local_stage + 41472), "n"(160) + : "memory"); + asm volatile("cp.async.bulk.commit_group;" ::: "memory"); + asm volatile("cp.async.bulk.wait_group 0;" ::: "memory"); + asm volatile("st.global.release.gpu.u32 [%0], %1;" + :: "l"(flag), "r"(ready_value) : "memory"); +} + +__device__ __forceinline__ void wait_k1_global_flag( + const unsigned int* flag, unsigned int expected) { + unsigned int ready = 0; + do { + asm volatile("ld.global.acquire.gpu.u32 %0, [%1];" + : "=r"(ready) : "l"(flag) : "memory"); + } while (ready != expected); +} + +__device__ __forceinline__ void wait_k1_global_ready( + const unsigned int* flag, unsigned int ready_value) { + unsigned int state = 0; + do { + asm volatile("ld.global.acquire.gpu.u32 %0, [%1];" + : "=r"(state) : "l"(flag) : "memory"); + } while ((state & ~6u) != ready_value); +} + +__device__ __forceinline__ void acknowledge_k1_global_packet( + unsigned int* flag, unsigned int owner_rank) { + unsigned int old; + const unsigned int mask = 1u << (owner_rank + 1u); + asm volatile("atom.global.release.gpu.or.b32 %0, [%1], %2;" + : "=r"(old) : "l"(flag), "r"(mask) : "memory"); +} + +__device__ __forceinline__ void load_k1_from_global( + uint32_t local_stage, uint32_t qk_barrier, + const unsigned char* packet) { + mbarrier_arrive_expect_tx(qk_barrier, kK1PacketBytes); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage), "l"(packet), "n"(28672), "r"(qk_barrier) + : "memory"); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage + 28672), "l"(packet + 28672), "n"(2688), + "r"(qk_barrier) : "memory"); + asm volatile( + "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" + " [%0], [%1], %2, [%3];" + :: "r"(local_stage + 41472), "l"(packet + 31360), "n"(160), + "r"(qk_barrier) : "memory"); +} + +extern "C" { + +__global__ __launch_bounds__(1024) void +// FLASHINFER INTEGRATION BEGIN: allow exact state alias +kernel_flashkda_bf16_fused_m64_k1_parallel(__nv_bfloat16* __restrict__ q, const void* __restrict__ q_tma, __nv_bfloat16* __restrict__ k, const void* __restrict__ k_tma, __nv_bfloat16* __restrict__ v, const void* __restrict__ v_tma, __nv_bfloat16* __restrict__ g, const void* __restrict__ g_tma, __nv_bfloat16* __restrict__ beta, const void* __restrict__ beta_tma, float* __restrict__ A_log, float* __restrict__ dt_bias, long long* __restrict__ cu_seqlens, int* __restrict__ seq_order, __nv_bfloat16* initial_state, __nv_bfloat16* __restrict__ out, const void* __restrict__ out_tma, __nv_bfloat16* final_state, unsigned char* k1_workspace, unsigned int* k1_flags, int mailbox_depth, int cluster_size, int num_heads, int use_initial_state, int store_final_state, float scale, float lower_bound) +// FLASHINFER INTEGRATION END: allow exact state alias +{ + // FLASHINFER INTEGRATION BEGIN: acquire global tensor maps + // CUDA kernel-start ordering does not acquire the tensor-map proxy. + // One thread acquires each 128-byte map; the CTA barrier publishes those + // acquires to every thread before any TMA instruction can use a map. + if (threadIdx.x == 0) { + asm volatile( + "fence.proxy.tensormap::generic.acquire.gpu [%0], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%1], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%2], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%3], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%4], 128;\n" + "fence.proxy.tensormap::generic.acquire.gpu [%5], 128;\n" + :: "l"(q_tma), "l"(k_tma), "l"(v_tma), "l"(g_tma), + "l"(beta_tma), "l"(out_tma) + : "memory"); + } + __syncthreads(); + // FLASHINFER INTEGRATION END: acquire global tensor maps + const int tid = threadIdx.x; + const int warp = make_warp_uniform(tid / 32); + const int lane = tid % 32; + const int kClusterSize = cluster_size; + constexpr int kOwnerCount = 2; + constexpr int kProducerFirstRank = kOwnerCount; + const int kProducerCount = kClusterSize - kOwnerCount; + + const int cta_rank = cluster_rank(); + const int bid = int(blockIdx.x) / kClusterSize; + + extern __shared__ __align__(1024) char smem_raw[]; + int smem; + smem = (int)(unsigned long long)__cvta_generic_to_shared(smem_raw); + + // Kernel setup ops + __nv_bfloat16* smem_qd = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_qd_addr = smem + 1024; + __nv_bfloat16* smem_g_raw = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_g_raw_addr = smem + 1024; + __nv_bfloat16* smem_g_raw_all = reinterpret_cast<__nv_bfloat16*>(smem_raw + 1024); + const int smem_g_raw_all_addr = smem + 1024; + __nv_bfloat16* smem_kd = reinterpret_cast<__nv_bfloat16*>(smem_raw + 9216); + const int smem_kd_addr = smem + 9216; + __nv_bfloat16* smem_q_raw_prefetch = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_q_raw_prefetch_addr = smem + 17408; + __nv_bfloat16* smem_final_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_final_trans_addr = smem + 17408; + __nv_bfloat16* smem_kr_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_kr_trans_addr = smem + 17408; + __nv_bfloat16* smem_mqk_trans = reinterpret_cast<__nv_bfloat16*>(smem_raw + 25600); + const int smem_mqk_trans_addr = smem + 25600; + __nv_bfloat16* smem_inv = reinterpret_cast<__nv_bfloat16*>(smem_raw + 29696); + const int smem_inv_addr = smem + 29696; + __nv_bfloat16* smem_v = reinterpret_cast<__nv_bfloat16*>(smem_raw + 32384); + const int smem_v_addr = smem + 32384; + __nv_bfloat16* smem_ki = reinterpret_cast<__nv_bfloat16*>(smem_raw + 17408); + const int smem_ki_addr = smem + 17408; + float* smem_gate = reinterpret_cast(smem_raw + 25600); + const int smem_gate_addr = smem + 25600; + __nv_bfloat16* smem_beta_raw = reinterpret_cast<__nv_bfloat16*>(smem_raw + 41984); + const int smem_beta_raw_addr = smem + 41984; + __nv_bfloat16* smem_inv_work = reinterpret_cast<__nv_bfloat16*>(smem_raw + 32384); + const int smem_inv_work_addr = smem + 32384; + __nv_bfloat16* smem_out = reinterpret_cast<__nv_bfloat16*>(smem_raw + 210944); + const int smem_out_addr = smem + 210944; + float* smem_restore_factor_all = reinterpret_cast(smem_raw + 41984); + const int smem_restore_factor_all_addr = smem + 41984; + float* smem_gt_prefix_all = reinterpret_cast(smem_raw + 41472); + const int smem_gt_prefix_all_addr = smem + 41472; + float* smem_gt_all = reinterpret_cast(smem_raw + 31744); + const int smem_gt_all_addr = smem + 31744; + float* smem_prep_beta_all = reinterpret_cast(smem_raw + 42500); + const int smem_prep_beta_all_addr = smem + 42500; + float* smem_gate_rate_all = reinterpret_cast(smem_raw + 42628); + const int smem_gate_rate_all_addr = smem + 42628; + float* smem_gate_all = reinterpret_cast(smem_raw + 25600); + const int smem_gate_all_addr = smem + 25600; + + // Mbarrier init (17 groups, 77 barriers) + // Mbarriers at smem_raw[0..616) + + if (warp == 0) { + uint32_t leader = elect_sync(); + // --- pipeline 'chunk_pipe' --- + // qk_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 0, 1, leader); + mbarrier_init_pred(smem + 8, 1, leader); + mbarrier_init_pred(smem + 16, 1, leader); + mbarrier_init_pred(smem + 24, 1, leader); + mbarrier_init_pred(smem + 32, 1, leader); + // gate_raw_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 40, 1, leader); + mbarrier_init_pred(smem + 48, 1, leader); + mbarrier_init_pred(smem + 56, 1, leader); + mbarrier_init_pred(smem + 64, 1, leader); + mbarrier_init_pred(smem + 72, 1, leader); + // qk_raw_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 80, 1, leader); + mbarrier_init_pred(smem + 88, 1, leader); + mbarrier_init_pred(smem + 96, 1, leader); + mbarrier_init_pred(smem + 104, 1, leader); + mbarrier_init_pred(smem + 112, 1, leader); + // v_full: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 120, 1, leader); + mbarrier_init_pred(smem + 128, 1, leader); + mbarrier_init_pred(smem + 136, 1, leader); + mbarrier_init_pred(smem + 144, 1, leader); + mbarrier_init_pred(smem + 152, 1, leader); + // v_free: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 160, 4, leader); + mbarrier_init_pred(smem + 168, 4, leader); + mbarrier_init_pred(smem + 176, 4, leader); + mbarrier_init_pred(smem + 184, 4, leader); + mbarrier_init_pred(smem + 192, 4, leader); + const int producer_free_arrivals = 1; + // smem_free: 5 barriers + mbarrier_init_pred(smem + 200, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 208, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 216, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 224, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 232, producer_free_arrivals, leader); + // raw_inputs_free: 5 barriers + mbarrier_init_pred(smem + 240, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 248, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 256, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 264, producer_free_arrivals, leader); + mbarrier_init_pred(smem + 272, producer_free_arrivals, leader); + // state_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 280, 4, leader); + mbarrier_init_pred(smem + 288, 4, leader); + mbarrier_init_pred(smem + 296, 4, leader); + mbarrier_init_pred(smem + 304, 4, leader); + mbarrier_init_pred(smem + 312, 4, leader); + // old_out_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 320, 1, leader); + mbarrier_init_pred(smem + 328, 1, leader); + mbarrier_init_pred(smem + 336, 1, leader); + mbarrier_init_pred(smem + 344, 1, leader); + mbarrier_init_pred(smem + 352, 1, leader); + // u_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 360, 4, leader); + mbarrier_init_pred(smem + 368, 4, leader); + mbarrier_init_pred(smem + 376, 4, leader); + mbarrier_init_pred(smem + 384, 4, leader); + mbarrier_init_pred(smem + 392, 4, leader); + // u2_acc_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 400, 1, leader); + mbarrier_init_pred(smem + 408, 1, leader); + mbarrier_init_pred(smem + 416, 1, leader); + mbarrier_init_pred(smem + 424, 1, leader); + mbarrier_init_pred(smem + 432, 1, leader); + // u2_inp_ready: 5 barriers, init_count=4 + mbarrier_init_pred(smem + 440, 4, leader); + mbarrier_init_pred(smem + 448, 4, leader); + mbarrier_init_pred(smem + 456, 4, leader); + mbarrier_init_pred(smem + 464, 4, leader); + mbarrier_init_pred(smem + 472, 4, leader); + // final_ready: 5 barriers, init_count=1 + mbarrier_init_pred(smem + 480, 1, leader); + mbarrier_init_pred(smem + 488, 1, leader); + mbarrier_init_pred(smem + 496, 1, leader); + mbarrier_init_pred(smem + 504, 1, leader); + mbarrier_init_pred(smem + 512, 1, leader); + // out_empty: 1 barriers, init_count=4 + mbarrier_init_pred(smem + 520, 4, leader); + // tmem_dealloc_ready: 1 barriers, init_count=2 + mbarrier_init_pred(smem + 528, 2, leader); + // prep_diag_ready: 5 barriers, init_count=2 + mbarrier_init_pred(smem + 536, 2, leader); + mbarrier_init_pred(smem + 544, 2, leader); + mbarrier_init_pred(smem + 552, 2, leader); + mbarrier_init_pred(smem + 560, 2, leader); + mbarrier_init_pred(smem + 568, 2, leader); + // prep_inv16_ready: 5 barriers, init_count=2 + mbarrier_init_pred(smem + 576, 2, leader); + mbarrier_init_pred(smem + 584, 2, leader); + mbarrier_init_pred(smem + 592, 2, leader); + mbarrier_init_pred(smem + 600, 2, leader); + mbarrier_init_pred(smem + 608, 2, leader); + asm volatile("fence.mbarrier_init.release.cluster;"); + } + + __syncthreads(); + cluster_sync(); + + // TMEM alloc (256 columns, 256 used) + volatile int* tmem_addr_storage = (volatile int*)(smem_raw + 656); + if (cta_rank < kOwnerCount && warp == 0) { + int _tmem_hold = smem + 656; + asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(_tmem_hold), "r"(256) : "memory"); + } else if (cta_rank >= kOwnerCount && tid == 0) { + tmem_addr_storage[0] = 0; + } + + __syncthreads(); + if (cta_rank < kOwnerCount) { + asm volatile("tcgen05.fence::after_thread_sync;"); + } + + const int mbar_base = smem; + #define qk_full_addr (mbar_base + 0) + #define gate_raw_full_addr (mbar_base + 40) + #define qk_raw_full_addr (mbar_base + 80) + #define v_full_addr (mbar_base + 120) + #define v_free_addr (mbar_base + 160) + #define smem_free_addr (mbar_base + 200) + #define raw_inputs_free_addr (mbar_base + 240) + #define state_inp_ready_addr (mbar_base + 280) + #define old_out_ready_addr (mbar_base + 320) + #define u_inp_ready_addr (mbar_base + 360) + #define u2_acc_ready_addr (mbar_base + 400) + #define u2_inp_ready_addr (mbar_base + 440) + #define final_ready_addr (mbar_base + 480) + #define out_empty_addr (mbar_base + 520) + #define tmem_dealloc_ready_addr (mbar_base + 528) + #define prep_diag_ready_addr (mbar_base + 536) + #define prep_inv16_ready_addr (mbar_base + 576) + const int taddr = tmem_addr_storage[0]; + + // Kernel post-init ops + const int tmem_tmem_state = taddr + 64; + const int tmem_tmem_state_inp = taddr; + const int tmem_tmem_u_acc = taddr + 224; + const int tmem_tmem_u2_inp = taddr + 224; + const int tmem_tmem_u2_acc = taddr; + const int tmem_tmem_out = taddr + 192; + const int tmem_tmem_state_out = taddr + 64; + + // ---- Register redistribution for WGs split across roles ---- + // Dec phase frees registers before any WG attempts inc. + if (cta_rank < kOwnerCount && warp >= 8) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + } + + // ---- Role: compute ---- + if (cta_rank < kOwnerCount && warp <= 3) { + asm volatile("setmaxnreg.inc.sync.aligned.u32 168;"); + { // compute_main + int task_idx = bid; + int value_row_offset = cta_rank * 64; + int seq_idx = seq_order[task_idx / num_heads]; + int head_idx = task_idx % num_heads; + long long bos = cu_seqlens[seq_idx]; + long long eos = cu_seqlens[seq_idx + 1]; + int seq_len = (int)(eos - bos); + int num_chunks = (seq_len + 32 - 1) / 32; + int warp_in_wg = warp % 4; + const int tmem_row_base = warp_in_wg * 32 << 16; + int lane_quad = lane & 3; + int local_row_top = warp_in_wg * 16 + lane / 4; + int local_row_bot = local_row_top + 8; + int state_row_top = value_row_offset + local_row_top; + int state_row_bot = value_row_offset + local_row_bot; + int warp_id_in_role = (warp - 0); + int compute_local_warp = warp_id_in_role; + long long state_head_base = ((long long)seq_idx * (long long)num_heads + (long long)head_idx) * 128 * 128; + long long state_base_top = state_head_base + (long long)state_row_top * 128; + long long state_base_bot = state_head_base + (long long)state_row_bot * 128; + #pragma unroll + for (int state_col_half = 0; state_col_half < 2; state_col_half++) { + float state_init[32]; + state_init[0] = 0.0f; + state_init[1] = 0.0f; + state_init[2] = 0.0f; + state_init[3] = 0.0f; + state_init[4] = 0.0f; + state_init[5] = 0.0f; + state_init[6] = 0.0f; + state_init[7] = 0.0f; + state_init[8] = 0.0f; + state_init[9] = 0.0f; + state_init[10] = 0.0f; + state_init[11] = 0.0f; + state_init[12] = 0.0f; + state_init[13] = 0.0f; + state_init[14] = 0.0f; + state_init[15] = 0.0f; + state_init[16] = 0.0f; + state_init[17] = 0.0f; + state_init[18] = 0.0f; + state_init[19] = 0.0f; + state_init[20] = 0.0f; + state_init[21] = 0.0f; + state_init[22] = 0.0f; + state_init[23] = 0.0f; + state_init[24] = 0.0f; + state_init[25] = 0.0f; + state_init[26] = 0.0f; + state_init[27] = 0.0f; + state_init[28] = 0.0f; + state_init[29] = 0.0f; + state_init[30] = 0.0f; + state_init[31] = 0.0f; + if (use_initial_state != 0) { + #pragma unroll + for (int state_col_group = 0; state_col_group < 8; state_col_group++) { + int state_col_pair = state_col_half * 64 + state_col_group * 8 + lane_quad * 2; + const int state_reg_base = state_col_group * 4; + state_init[state_reg_base] = (float)initial_state[state_base_top + (long long)state_col_pair]; + state_init[state_reg_base + 1] = (float)initial_state[state_base_top + (long long)state_col_pair + 1]; + state_init[state_reg_base + 2] = (float)initial_state[state_base_bot + (long long)state_col_pair]; + state_init[state_reg_base + 3] = (float)initial_state[state_base_bot + (long long)state_col_pair + 1]; + } + } + asm volatile( + "tcgen05.st.sync.aligned.16x256b.x8.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32};" + :: "r"(taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_half * 64)), "r"(*reinterpret_cast(&state_init[0])), "r"(*reinterpret_cast(&state_init[1])), "r"(*reinterpret_cast(&state_init[2])), "r"(*reinterpret_cast(&state_init[3])), "r"(*reinterpret_cast(&state_init[4])), "r"(*reinterpret_cast(&state_init[5])), "r"(*reinterpret_cast(&state_init[6])), "r"(*reinterpret_cast(&state_init[7])), "r"(*reinterpret_cast(&state_init[8])), "r"(*reinterpret_cast(&state_init[9])), "r"(*reinterpret_cast(&state_init[10])), "r"(*reinterpret_cast(&state_init[11])), "r"(*reinterpret_cast(&state_init[12])), "r"(*reinterpret_cast(&state_init[13])), "r"(*reinterpret_cast(&state_init[14])), "r"(*reinterpret_cast(&state_init[15])), "r"(*reinterpret_cast(&state_init[16])), "r"(*reinterpret_cast(&state_init[17])), "r"(*reinterpret_cast(&state_init[18])), "r"(*reinterpret_cast(&state_init[19])), "r"(*reinterpret_cast(&state_init[20])), "r"(*reinterpret_cast(&state_init[21])), "r"(*reinterpret_cast(&state_init[22])), "r"(*reinterpret_cast(&state_init[23])), "r"(*reinterpret_cast(&state_init[24])), "r"(*reinterpret_cast(&state_init[25])), "r"(*reinterpret_cast(&state_init[26])), "r"(*reinterpret_cast(&state_init[27])), "r"(*reinterpret_cast(&state_init[28])), "r"(*reinterpret_cast(&state_init[29])), "r"(*reinterpret_cast(&state_init[30])), "r"(*reinterpret_cast(&state_init[31])) + : "memory"); + } + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + unsigned int compute_stage = 0; + unsigned int _phase_qk_full = 0; + unsigned int _phase_v_full = 0; + unsigned int _phase_old_out_ready = 0; + unsigned int _phase_u2_acc_ready = 0; + unsigned int _phase_final_ready = 0; + #pragma unroll 1 + for (int chunk_idx = 0; chunk_idx < num_chunks; chunk_idx++) { + mbarrier_wait_cluster(qk_full_addr + (compute_stage) * 8, _phase_qk_full); + #pragma unroll + for (int state_col_half_1 = 0; state_col_half_1 < 2; state_col_half_1++) { + int state_addr = taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_half_1 * 64); + float _tmem_load_0[32]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x8.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, [%32];" + : "=r"(*reinterpret_cast(&_tmem_load_0[0])), "=r"(*reinterpret_cast(&_tmem_load_0[1])), "=r"(*reinterpret_cast(&_tmem_load_0[2])), "=r"(*reinterpret_cast(&_tmem_load_0[3])), "=r"(*reinterpret_cast(&_tmem_load_0[4])), "=r"(*reinterpret_cast(&_tmem_load_0[5])), "=r"(*reinterpret_cast(&_tmem_load_0[6])), "=r"(*reinterpret_cast(&_tmem_load_0[7])), "=r"(*reinterpret_cast(&_tmem_load_0[8])), "=r"(*reinterpret_cast(&_tmem_load_0[9])), "=r"(*reinterpret_cast(&_tmem_load_0[10])), "=r"(*reinterpret_cast(&_tmem_load_0[11])), "=r"(*reinterpret_cast(&_tmem_load_0[12])), "=r"(*reinterpret_cast(&_tmem_load_0[13])), "=r"(*reinterpret_cast(&_tmem_load_0[14])), "=r"(*reinterpret_cast(&_tmem_load_0[15])), "=r"(*reinterpret_cast(&_tmem_load_0[16])), "=r"(*reinterpret_cast(&_tmem_load_0[17])), "=r"(*reinterpret_cast(&_tmem_load_0[18])), "=r"(*reinterpret_cast(&_tmem_load_0[19])), "=r"(*reinterpret_cast(&_tmem_load_0[20])), "=r"(*reinterpret_cast(&_tmem_load_0[21])), "=r"(*reinterpret_cast(&_tmem_load_0[22])), "=r"(*reinterpret_cast(&_tmem_load_0[23])), "=r"(*reinterpret_cast(&_tmem_load_0[24])), "=r"(*reinterpret_cast(&_tmem_load_0[25])), "=r"(*reinterpret_cast(&_tmem_load_0[26])), "=r"(*reinterpret_cast(&_tmem_load_0[27])), "=r"(*reinterpret_cast(&_tmem_load_0[28])), "=r"(*reinterpret_cast(&_tmem_load_0[29])), "=r"(*reinterpret_cast(&_tmem_load_0[30])), "=r"(*reinterpret_cast(&_tmem_load_0[31])) + : "r"(state_addr) + : "memory"); + uint32_t _tmem_load_0_bf16[16]; + #pragma unroll + for (int _lp = 0; _lp < 16; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(_tmem_load_0[_lp*2 + 0], _tmem_load_0[_lp*2+1 + 0])); + _tmem_load_0_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile( + "tcgen05.st.sync.aligned.16x128b.x8.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16};" + :: "r"(taddr + (unsigned int)tmem_row_base + (unsigned int)(state_col_half_1 * 32)), "r"(*reinterpret_cast(&_tmem_load_0_bf16[0])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[1])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[2])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[3])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[4])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[5])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[6])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[7])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[8])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[9])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[10])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[11])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[12])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[13])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[14])), "r"(*reinterpret_cast(&_tmem_load_0_bf16[15])) + : "memory"); + #pragma unroll + for (int state_col_group_1 = 0; state_col_group_1 < 8; state_col_group_1++) { + int state_col_pair_1 = state_col_half_1 * 64 + state_col_group_1 * 8 + lane_quad * 2; + const int state_reg_base_1 = state_col_group_1 * 4; + float state_scale_0 = smem_gt_all[compute_stage * 10496 + (unsigned int)state_col_pair_1]; + float state_scale_1 = smem_gt_all[compute_stage * 10496 + (unsigned int)state_col_pair_1 + 1]; + _tmem_load_0[state_reg_base_1] = _tmem_load_0[state_reg_base_1] * state_scale_0; + _tmem_load_0[state_reg_base_1 + 1] = _tmem_load_0[state_reg_base_1 + 1] * state_scale_1; + _tmem_load_0[state_reg_base_1 + 2] = _tmem_load_0[state_reg_base_1 + 2] * state_scale_0; + _tmem_load_0[state_reg_base_1 + 3] = _tmem_load_0[state_reg_base_1 + 3] * state_scale_1; + } + asm volatile( + "tcgen05.st.sync.aligned.16x256b.x8.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32};" + :: "r"(state_addr), "r"(*reinterpret_cast(&_tmem_load_0[0])), "r"(*reinterpret_cast(&_tmem_load_0[1])), "r"(*reinterpret_cast(&_tmem_load_0[2])), "r"(*reinterpret_cast(&_tmem_load_0[3])), "r"(*reinterpret_cast(&_tmem_load_0[4])), "r"(*reinterpret_cast(&_tmem_load_0[5])), "r"(*reinterpret_cast(&_tmem_load_0[6])), "r"(*reinterpret_cast(&_tmem_load_0[7])), "r"(*reinterpret_cast(&_tmem_load_0[8])), "r"(*reinterpret_cast(&_tmem_load_0[9])), "r"(*reinterpret_cast(&_tmem_load_0[10])), "r"(*reinterpret_cast(&_tmem_load_0[11])), "r"(*reinterpret_cast(&_tmem_load_0[12])), "r"(*reinterpret_cast(&_tmem_load_0[13])), "r"(*reinterpret_cast(&_tmem_load_0[14])), "r"(*reinterpret_cast(&_tmem_load_0[15])), "r"(*reinterpret_cast(&_tmem_load_0[16])), "r"(*reinterpret_cast(&_tmem_load_0[17])), "r"(*reinterpret_cast(&_tmem_load_0[18])), "r"(*reinterpret_cast(&_tmem_load_0[19])), "r"(*reinterpret_cast(&_tmem_load_0[20])), "r"(*reinterpret_cast(&_tmem_load_0[21])), "r"(*reinterpret_cast(&_tmem_load_0[22])), "r"(*reinterpret_cast(&_tmem_load_0[23])), "r"(*reinterpret_cast(&_tmem_load_0[24])), "r"(*reinterpret_cast(&_tmem_load_0[25])), "r"(*reinterpret_cast(&_tmem_load_0[26])), "r"(*reinterpret_cast(&_tmem_load_0[27])), "r"(*reinterpret_cast(&_tmem_load_0[28])), "r"(*reinterpret_cast(&_tmem_load_0[29])), "r"(*reinterpret_cast(&_tmem_load_0[30])), "r"(*reinterpret_cast(&_tmem_load_0[31])) + : "memory"); + } + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(state_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(v_full_addr + (compute_stage) * 8, _phase_v_full); + mbarrier_wait(old_out_ready_addr + (compute_stage) * 8, _phase_old_out_ready); + float _tmem_load_1[16]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x4.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=r"(*reinterpret_cast(&_tmem_load_1[0])), "=r"(*reinterpret_cast(&_tmem_load_1[1])), "=r"(*reinterpret_cast(&_tmem_load_1[2])), "=r"(*reinterpret_cast(&_tmem_load_1[3])), "=r"(*reinterpret_cast(&_tmem_load_1[4])), "=r"(*reinterpret_cast(&_tmem_load_1[5])), "=r"(*reinterpret_cast(&_tmem_load_1[6])), "=r"(*reinterpret_cast(&_tmem_load_1[7])), "=r"(*reinterpret_cast(&_tmem_load_1[8])), "=r"(*reinterpret_cast(&_tmem_load_1[9])), "=r"(*reinterpret_cast(&_tmem_load_1[10])), "=r"(*reinterpret_cast(&_tmem_load_1[11])), "=r"(*reinterpret_cast(&_tmem_load_1[12])), "=r"(*reinterpret_cast(&_tmem_load_1[13])), "=r"(*reinterpret_cast(&_tmem_load_1[14])), "=r"(*reinterpret_cast(&_tmem_load_1[15])) + : "r"(taddr + 224 + (unsigned int)tmem_row_base) + : "memory"); + float residual_values[16]; + int v_stage_addr = smem_v_addr + compute_stage * 41984; + unsigned int v_ld_bits[2]; + #pragma unroll + for (int token_group = 0; token_group < 4; token_group++) { + int token_pair = token_group * 8 + lane_quad * 2; + const int residual_reg_base = token_group * 4; + float beta_0 = smem_prep_beta_all[compute_stage * 10496 + (unsigned int)token_pair]; + float beta_1 = smem_prep_beta_all[compute_stage * 10496 + (unsigned int)token_pair + 1]; + int v_ld_matrix = lane / 8 & 1; + int v_ld_token = token_group * 8 + (lane & 7); + int v_ld_row = warp_in_wg * 16 + v_ld_matrix * 8; + int v_ld_row_addr = v_stage_addr + v_ld_token * 64 * 2; + int v_ld_addr = (v_ld_row_addr + (v_ld_row * 2 ^ (v_ld_row_addr >> 7 & 7) << 4)); + asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0, %1}, [%2];\n" + : "=r"(v_ld_bits[0]), "=r"(v_ld_bits[1]) + : "r"(v_ld_addr) + : "memory"); + float v_ld_bits_fp32[4]; + #pragma unroll + for (int _pair = 0; _pair < 2; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&v_ld_bits_fp32[_pair * 2])[0]), "=f"((&v_ld_bits_fp32[_pair * 2])[1]) + : "r"(v_ld_bits[_pair + 0])); + } + residual_values[residual_reg_base] = (v_ld_bits_fp32[0] - _tmem_load_1[residual_reg_base]) * beta_0; + residual_values[residual_reg_base + 1] = (v_ld_bits_fp32[1] - _tmem_load_1[residual_reg_base + 1]) * beta_1; + residual_values[residual_reg_base + 2] = (v_ld_bits_fp32[2] - _tmem_load_1[residual_reg_base + 2]) * beta_0; + residual_values[residual_reg_base + 3] = (v_ld_bits_fp32[3] - _tmem_load_1[residual_reg_base + 3]) * beta_1; + } + uint32_t residual_values_bf16[8]; + #pragma unroll + for (int _lp = 0; _lp < 8; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(residual_values[_lp*2 + 0], residual_values[_lp*2+1 + 0])); + residual_values_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile( + "tcgen05.st.sync.aligned.16x128b.x4.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8};" + :: "r"(taddr + 224 + (unsigned int)tmem_row_base), "r"(*reinterpret_cast(&residual_values_bf16[0])), "r"(*reinterpret_cast(&residual_values_bf16[1])), "r"(*reinterpret_cast(&residual_values_bf16[2])), "r"(*reinterpret_cast(&residual_values_bf16[3])), "r"(*reinterpret_cast(&residual_values_bf16[4])), "r"(*reinterpret_cast(&residual_values_bf16[5])), "r"(*reinterpret_cast(&residual_values_bf16[6])), "r"(*reinterpret_cast(&residual_values_bf16[7])) + : "memory"); + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(v_free_addr + (compute_stage) * 8); + mbarrier_arrive(u_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(u2_acc_ready_addr + (compute_stage) * 8, _phase_u2_acc_ready); + float _tmem_load_2[16]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x4.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=r"(*reinterpret_cast(&_tmem_load_2[0])), "=r"(*reinterpret_cast(&_tmem_load_2[1])), "=r"(*reinterpret_cast(&_tmem_load_2[2])), "=r"(*reinterpret_cast(&_tmem_load_2[3])), "=r"(*reinterpret_cast(&_tmem_load_2[4])), "=r"(*reinterpret_cast(&_tmem_load_2[5])), "=r"(*reinterpret_cast(&_tmem_load_2[6])), "=r"(*reinterpret_cast(&_tmem_load_2[7])), "=r"(*reinterpret_cast(&_tmem_load_2[8])), "=r"(*reinterpret_cast(&_tmem_load_2[9])), "=r"(*reinterpret_cast(&_tmem_load_2[10])), "=r"(*reinterpret_cast(&_tmem_load_2[11])), "=r"(*reinterpret_cast(&_tmem_load_2[12])), "=r"(*reinterpret_cast(&_tmem_load_2[13])), "=r"(*reinterpret_cast(&_tmem_load_2[14])), "=r"(*reinterpret_cast(&_tmem_load_2[15])) + : "r"(taddr + (unsigned int)tmem_row_base) + : "memory"); + uint32_t _tmem_load_2_bf16[8]; + #pragma unroll + for (int _lp = 0; _lp < 8; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(_tmem_load_2[_lp*2 + 0], _tmem_load_2[_lp*2+1 + 0])); + _tmem_load_2_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile( + "tcgen05.st.sync.aligned.16x128b.x4.b32" + " [%0], {%1, %2, %3, %4, %5, %6, %7, %8};" + :: "r"(taddr + 224 + (unsigned int)tmem_row_base), "r"(*reinterpret_cast(&_tmem_load_2_bf16[0])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[1])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[2])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[3])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[4])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[5])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[6])), "r"(*reinterpret_cast(&_tmem_load_2_bf16[7])) + : "memory"); + asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(u2_inp_ready_addr + (compute_stage) * 8); + } + mbarrier_wait(final_ready_addr + (compute_stage) * 8, _phase_final_ready); + compute_stage += 1; + if (compute_stage == 5) { compute_stage = 0; _phase_qk_full ^= 1; _phase_v_full ^= 1; _phase_old_out_ready ^= 1; _phase_u2_acc_ready ^= 1; _phase_final_ready ^= 1; } + } + if (store_final_state != 0) { + #pragma unroll + for (int state_col_half_2 = 0; state_col_half_2 < 2; state_col_half_2++) { + float _tmem_load_3[32]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x8.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, [%32];" + : "=r"(*reinterpret_cast(&_tmem_load_3[0])), "=r"(*reinterpret_cast(&_tmem_load_3[1])), "=r"(*reinterpret_cast(&_tmem_load_3[2])), "=r"(*reinterpret_cast(&_tmem_load_3[3])), "=r"(*reinterpret_cast(&_tmem_load_3[4])), "=r"(*reinterpret_cast(&_tmem_load_3[5])), "=r"(*reinterpret_cast(&_tmem_load_3[6])), "=r"(*reinterpret_cast(&_tmem_load_3[7])), "=r"(*reinterpret_cast(&_tmem_load_3[8])), "=r"(*reinterpret_cast(&_tmem_load_3[9])), "=r"(*reinterpret_cast(&_tmem_load_3[10])), "=r"(*reinterpret_cast(&_tmem_load_3[11])), "=r"(*reinterpret_cast(&_tmem_load_3[12])), "=r"(*reinterpret_cast(&_tmem_load_3[13])), "=r"(*reinterpret_cast(&_tmem_load_3[14])), "=r"(*reinterpret_cast(&_tmem_load_3[15])), "=r"(*reinterpret_cast(&_tmem_load_3[16])), "=r"(*reinterpret_cast(&_tmem_load_3[17])), "=r"(*reinterpret_cast(&_tmem_load_3[18])), "=r"(*reinterpret_cast(&_tmem_load_3[19])), "=r"(*reinterpret_cast(&_tmem_load_3[20])), "=r"(*reinterpret_cast(&_tmem_load_3[21])), "=r"(*reinterpret_cast(&_tmem_load_3[22])), "=r"(*reinterpret_cast(&_tmem_load_3[23])), "=r"(*reinterpret_cast(&_tmem_load_3[24])), "=r"(*reinterpret_cast(&_tmem_load_3[25])), "=r"(*reinterpret_cast(&_tmem_load_3[26])), "=r"(*reinterpret_cast(&_tmem_load_3[27])), "=r"(*reinterpret_cast(&_tmem_load_3[28])), "=r"(*reinterpret_cast(&_tmem_load_3[29])), "=r"(*reinterpret_cast(&_tmem_load_3[30])), "=r"(*reinterpret_cast(&_tmem_load_3[31])) + : "r"(taddr + 64 + (unsigned int)tmem_row_base + (unsigned int)(state_col_half_2 * 64)) + : "memory"); + #pragma unroll + for (int state_col_group_2 = 0; state_col_group_2 < 8; state_col_group_2++) { + int state_col_pair_2 = state_col_half_2 * 64 + state_col_group_2 * 8 + lane_quad * 2; + const int state_reg_base_2 = state_col_group_2 * 4; + final_state[state_base_top + (long long)state_col_pair_2] = _tmem_load_3[state_reg_base_2]; + final_state[state_base_top + (long long)state_col_pair_2 + 1] = _tmem_load_3[state_reg_base_2 + 1]; + final_state[state_base_bot + (long long)state_col_pair_2] = _tmem_load_3[state_reg_base_2 + 2]; + final_state[state_base_bot + (long long)state_col_pair_2 + 1] = _tmem_load_3[state_reg_base_2 + 3]; + } + } + } + asm volatile("barrier.sync 10, 128;" ::: "memory"); + if (compute_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive(tmem_dealloc_ready_addr); + } + } + } + // ---- Role: epilogue ---- + } else if (cta_rank < kOwnerCount && warp >= 4 && warp <= 7) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + { // epilogue_main + int task_idx_1 = bid; + int value_split_idx_1 = cta_rank; + int value_row_offset_1 = value_split_idx_1 * 64; + int seq_idx_1 = seq_order[task_idx_1 / num_heads]; + int head_idx_1 = task_idx_1 % num_heads; + long long bos_1 = cu_seqlens[seq_idx_1]; + long long eos_1 = cu_seqlens[seq_idx_1 + 1]; + int seq_len_1 = (int)(eos_1 - bos_1); + int num_chunks_1 = (seq_len_1 + 32 - 1) / 32; + int warp_id_in_role_1 = (warp - 4); + int epilogue_local_warp = warp_id_in_role_1; + int warp_in_wg_1 = warp % 4; + const int tmem_row_base_1 = warp_in_wg_1 * 32 << 16; + int lane_quad_1 = lane & 3; + int local_row_top_1 = warp_in_wg_1 * 16 + lane / 4; + int local_row_bot_1 = local_row_top_1 + 8; + int state_row_top_1 = value_row_offset_1 + local_row_top_1; + int state_row_bot_1 = value_row_offset_1 + local_row_bot_1; + unsigned int epilogue_stage = 0; + unsigned int output_stage = 0; + unsigned int _phase_final_ready_1 = 0; + #pragma unroll 1 + for (int chunk_idx_1 = 0; chunk_idx_1 < num_chunks_1; chunk_idx_1++) { + mbarrier_wait(final_ready_addr + (epilogue_stage) * 8, _phase_final_ready_1); + int chunk_is_full = ((seq_len_1 >= (chunk_idx_1 + 1) * 32) ? 1 : 0); + if (chunk_is_full != 0) { + float _tmem_load_4[16]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x4.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=r"(*reinterpret_cast(&_tmem_load_4[0])), "=r"(*reinterpret_cast(&_tmem_load_4[1])), "=r"(*reinterpret_cast(&_tmem_load_4[2])), "=r"(*reinterpret_cast(&_tmem_load_4[3])), "=r"(*reinterpret_cast(&_tmem_load_4[4])), "=r"(*reinterpret_cast(&_tmem_load_4[5])), "=r"(*reinterpret_cast(&_tmem_load_4[6])), "=r"(*reinterpret_cast(&_tmem_load_4[7])), "=r"(*reinterpret_cast(&_tmem_load_4[8])), "=r"(*reinterpret_cast(&_tmem_load_4[9])), "=r"(*reinterpret_cast(&_tmem_load_4[10])), "=r"(*reinterpret_cast(&_tmem_load_4[11])), "=r"(*reinterpret_cast(&_tmem_load_4[12])), "=r"(*reinterpret_cast(&_tmem_load_4[13])), "=r"(*reinterpret_cast(&_tmem_load_4[14])), "=r"(*reinterpret_cast(&_tmem_load_4[15])) + : "r"(taddr + 192 + (unsigned int)tmem_row_base_1) + : "memory"); + asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(out_empty_addr); + } + if (epilogue_local_warp == 0) { + if (chunk_idx_1 >= 2) { + asm volatile("cp.async.bulk.wait_group.read 1;"); + } + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + int out_stage_addr = smem_out_addr + output_stage * 4096; + unsigned int out_packed[8]; + #pragma unroll + for (int _lp = 0; _lp < 8; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(_tmem_load_4[_lp*2 + 0], _tmem_load_4[_lp*2+1 + 0])); + out_packed[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int token_group_1 = 0; token_group_1 < 2; token_group_1++) { + int mtx_idx = lane / 8; + int row_addr = lane & 7; + int dim_base = epilogue_local_warp * 16 + (mtx_idx & 1) * 8; + int token_base = token_group_1 * 16 + mtx_idx / 2 * 8; + int token_addr = token_base + row_addr; + int token_pair_1 = token_addr / 2; + int token_parity = token_addr & 1; + int raw_row = token_pair_1; + int raw_col = (dim_base & 63 ^ (token_pair_1 & 3) << 4 ^ token_parity << 3) + token_parity * 64; + int stsm_offset = (raw_row * 128 + raw_col) * 2; + const int pack_base = token_group_1 * 4; + uint32_t _stmatrix_addr_0 = static_cast((unsigned long long)(out_stage_addr + stsm_offset)); + asm volatile("stmatrix.sync.aligned.m8n8.x4.trans.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_0), "r"(*reinterpret_cast(&out_packed[pack_base])), "r"(*reinterpret_cast(&out_packed[pack_base + 1])), "r"(*reinterpret_cast(&out_packed[pack_base + 2])), "r"(*reinterpret_cast(&out_packed[pack_base + 3])) + : "memory"); + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0) { + if (elect_sync()) { + tma_store_4d(out_tma, 0, (int)(bos_1 + (long long)(chunk_idx_1 * 32)), head_idx_1, value_split_idx_1, out_stage_addr); + } + asm volatile("cp.async.bulk.commit_group;"); + } + output_stage = output_stage ^ 1; + } else { + float _tmem_load_5[16]; + asm volatile( + "tcgen05.ld.sync.aligned.16x256b.x4.b32" + " {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];" + : "=r"(*reinterpret_cast(&_tmem_load_5[0])), "=r"(*reinterpret_cast(&_tmem_load_5[1])), "=r"(*reinterpret_cast(&_tmem_load_5[2])), "=r"(*reinterpret_cast(&_tmem_load_5[3])), "=r"(*reinterpret_cast(&_tmem_load_5[4])), "=r"(*reinterpret_cast(&_tmem_load_5[5])), "=r"(*reinterpret_cast(&_tmem_load_5[6])), "=r"(*reinterpret_cast(&_tmem_load_5[7])), "=r"(*reinterpret_cast(&_tmem_load_5[8])), "=r"(*reinterpret_cast(&_tmem_load_5[9])), "=r"(*reinterpret_cast(&_tmem_load_5[10])), "=r"(*reinterpret_cast(&_tmem_load_5[11])), "=r"(*reinterpret_cast(&_tmem_load_5[12])), "=r"(*reinterpret_cast(&_tmem_load_5[13])), "=r"(*reinterpret_cast(&_tmem_load_5[14])), "=r"(*reinterpret_cast(&_tmem_load_5[15])) + : "r"(taddr + 192 + (unsigned int)tmem_row_base_1) + : "memory"); + asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory"); + if (elect_sync()) { + mbarrier_arrive(out_empty_addr); + } + #pragma unroll + for (int token_group_2 = 0; token_group_2 < 4; token_group_2++) { + int token_pair_2 = token_group_2 * 8 + lane_quad_1 * 2; + const int out_reg_base = token_group_2 * 4; + long long out_token_0 = bos_1 + (long long)(chunk_idx_1 * 32 + token_pair_2); + long long out_token_1 = out_token_0 + 1; + if (out_token_0 < eos_1) { + long long out_idx_top_0 = (out_token_0 * (long long)num_heads + (long long)head_idx_1) * 128 + (long long)state_row_top_1; + long long out_idx_bot_0 = (out_token_0 * (long long)num_heads + (long long)head_idx_1) * 128 + (long long)state_row_bot_1; + out[out_idx_top_0] = _tmem_load_5[out_reg_base]; + out[out_idx_bot_0] = _tmem_load_5[out_reg_base + 2]; + } + if (out_token_1 < eos_1) { + long long out_idx_top_1 = (out_token_1 * (long long)num_heads + (long long)head_idx_1) * 128 + (long long)state_row_top_1; + long long out_idx_bot_1 = (out_token_1 * (long long)num_heads + (long long)head_idx_1) * 128 + (long long)state_row_bot_1; + out[out_idx_top_1] = _tmem_load_5[out_reg_base + 1]; + out[out_idx_bot_1] = _tmem_load_5[out_reg_base + 3]; + } + } + } + if (epilogue_local_warp == 0 && elect_sync()) { + long long packet_idx = + (long long)task_idx_1 * mailbox_depth + + chunk_idx_1 % mailbox_depth; + acknowledge_k1_global_packet( + k1_flags + packet_idx, (unsigned int)cta_rank); + mbarrier_arrive( + raw_inputs_free_addr + epilogue_stage * 8); + mbarrier_arrive(smem_free_addr + epilogue_stage * 8); + } + epilogue_stage += 1; + if (epilogue_stage == 5) { epilogue_stage = 0; _phase_final_ready_1 ^= 1; } + } + if (epilogue_local_warp == 0) { + asm volatile("cp.async.bulk.wait_group 0;"); + } + asm volatile("barrier.sync 9, 128;" ::: "memory"); + if (epilogue_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive(tmem_dealloc_ready_addr); + } + } + } + // ---- Role: mma ---- + } else if (cta_rank < kOwnerCount && warp == 9) { + { // mma_main + int task_idx_2 = bid; + int seq_idx_2 = seq_order[task_idx_2 / num_heads]; + long long bos_2 = cu_seqlens[seq_idx_2]; + long long eos_2 = cu_seqlens[seq_idx_2 + 1]; + int seq_len_2 = (int)(eos_2 - bos_2); + int num_chunks_2 = (seq_len_2 + 32 - 1) / 32; + unsigned int mma_stage = 0; + unsigned int _phase_qk_full_1 = 0; + unsigned int _phase_state_inp_ready = 0; + unsigned int _phase_out_empty_0 = 1; + unsigned int _phase_u_inp_ready = 0; + unsigned int _phase_u2_inp_ready = 0; + #pragma unroll 1 + for (int _chunk_idx = 0; _chunk_idx < num_chunks_2; _chunk_idx++) { + mbarrier_wait_cluster(qk_full_addr + (mma_stage) * 8, _phase_qk_full_1); + mbarrier_wait(state_inp_ready_addr + (mma_stage) * 8, _phase_state_inp_ready); + mbarrier_wait(out_empty_addr, _phase_out_empty_0); + _phase_out_empty_0 ^= 1; + int _mma_b_addr_0 = smem_qd_addr + mma_stage * 41984; + int _mma_b_lo_0 = make_warp_uniform((_mma_b_addr_0 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 67634320;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 16], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 24], db, id, p1;\n\t" + "add.u32 blo, blo, 250;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 32], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 40], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 48], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 56], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_out), "r"(_mma_b_lo_0), "r"(tmem_tmem_state_inp), "r"(0)); + int _mma_b_addr_1 = smem_kd_addr + mma_stage * 41984; + int _mma_b_lo_1 = make_warp_uniform((_mma_b_addr_1 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 67634320;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 16], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 24], db, id, p1;\n\t" + "add.u32 blo, blo, 250;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 32], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 40], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 48], db, id, p1;\n\t" + "add.u32 blo, blo, 2;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 56], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_u_acc), "r"(_mma_b_lo_1), "r"(tmem_tmem_state_inp), "r"(0)); + elect_commit(old_out_ready_addr + mma_stage * 8); + mbarrier_wait(u_inp_ready_addr + (mma_stage) * 8, _phase_u_inp_ready); + int _mma_b_addr_2 = smem_inv_addr + mma_stage * 41984; + int _mma_b_lo_2 = make_warp_uniform((_mma_b_addr_2 >> 4) & 0x3FFF); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0xC0004010;\n\t" + "mov.b32 id, 67634320;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 64;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_u2_acc), "r"(_mma_b_lo_2), "r"(tmem_tmem_u2_inp), "r"(0)); + elect_commit(u2_acc_ready_addr + (mma_stage) * 8); + mbarrier_wait(u2_inp_ready_addr + (mma_stage) * 8, _phase_u2_inp_ready); + int _mma_b_addr_3 = smem_final_trans_addr + mma_stage * 41984; + int _mma_b_lo_3 = make_warp_uniform(((_mma_b_addr_3 >> 4) & 0x3FFF) | 0x1000000); + asm volatile( + "{\n\t" + ".reg .pred leader, p0, p1;\n\t" + ".reg .b32 dhi, blo, id;\n\t" + ".reg .b64 db;\n\t" + "elect.sync _|leader, 0xFFFFFFFF;\n\t" + "setp.ne.b32 p0, %3, 0;\n\t" + "setp.ne.b32 p1, 1, 0;\n\t" + "" + "mov.b32 dhi, 0x40004040;\n\t" + "mov.b32 id, 69797008;\n\t" + "mov.b32 blo, %1;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2], db, id, p0;\n\t" + "add.u32 blo, blo, 128;\n\t" + "mov.b64 db, {blo, dhi};\n\t" + "@leader tcgen05.mma.cta_group::1.kind::f16 [%0], [%2 + 8], db, id, p1;\n\t" + "}\n" + :: "r"(tmem_tmem_state_out), "r"(_mma_b_lo_3), "r"(tmem_tmem_u2_inp), "r"(1)); + elect_commit(final_ready_addr + mma_stage * 8); + mma_stage += 1; + if (mma_stage == 5) { mma_stage = 0; _phase_qk_full_1 ^= 1; _phase_state_inp_ready ^= 1; _phase_u_inp_ready ^= 1; _phase_u2_inp_ready ^= 1; } + } + unsigned int _phase_tmem_dealloc_ready_0 = 0; + mbarrier_wait(tmem_dealloc_ready_addr, _phase_tmem_dealloc_ready_0); + _phase_tmem_dealloc_ready_0 ^= 1; + int _tmem_dealloc_addr = *((volatile int*)tmem_addr_storage); + asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(_tmem_dealloc_addr), "r"(256)); + asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;"); + } + // ---- Role: load ---- + } else if (cta_rank < kOwnerCount && warp == 10) { + { // load_main + int task_idx_3 = bid; + int value_row_offset_2 = cta_rank * 64; + int seq_idx_3 = seq_order[task_idx_3 / num_heads]; + int head_idx_2 = task_idx_3 % num_heads; + long long bos_3 = cu_seqlens[seq_idx_3]; + long long eos_3 = cu_seqlens[seq_idx_3 + 1]; + int seq_len_3 = (int)(eos_3 - bos_3); + int num_chunks_3 = (seq_len_3 + 32 - 1) / 32; + unsigned int load_stage = 0; + unsigned int _phase_v_free = 1; + unsigned int _phase_qk_full_2 = 0; + #pragma unroll 1 + for (int chunk_idx_2 = 0; chunk_idx_2 < num_chunks_3; chunk_idx_2++) { + mbarrier_wait(v_free_addr + (load_stage) * 8, _phase_v_free); + mbarrier_wait_cluster(qk_full_addr + (load_stage) * 8, _phase_qk_full_2); + int chunk_is_full_1 = ((seq_len_3 >= (chunk_idx_2 + 1) * 32) ? 1 : 0); + if (elect_sync()) { + if (chunk_is_full_1 != 0) { + mbarrier_arrive_expect_tx(v_full_addr + (load_stage) * 8, 4096); + tma_3d_gmem2smem(smem_v_addr + load_stage * 41984, v_tma, value_row_offset_2, head_idx_2, (int)(bos_3 + (long long)(chunk_idx_2 * 32)), v_full_addr + (load_stage) * 8); + } + } + if (chunk_is_full_1 == 0) { + #pragma unroll + for (int v_load_iter = 0; v_load_iter < 8; v_load_iter++) { + int v_item = v_load_iter * 32 + lane; + int row = v_item / 8; + int segment = v_item % 8; + long long token = bos_3 + (long long)(chunk_idx_2 * 32 + row); + int token_valid = ((token < eos_3) ? 1 : 0); + long long v_src = (token * (long long)num_heads + (long long)head_idx_2) * 128 + (long long)value_row_offset_2 + (long long)(segment * 8); + int v_dst_row_addr = smem_v_addr + load_stage * 41984 + (unsigned int)(row * 64 * 2); + int v_dst_addr = (v_dst_row_addr + (segment * 8 * 2 ^ (v_dst_row_addr >> 7 & 7) << 4)); + asm volatile("cp.async.cg.shared::cta.global [%0], [%1], 16, %2;" + :: "r"(v_dst_addr), "l"(v + v_src), "r"((token_valid != 0) ? 16 : 0)); + } + asm volatile("cp.async.commit_group;"); + asm volatile("cp.async.wait_group 0;"); + } + asm volatile("barrier.sync 8, 32;" ::: "memory"); + if (elect_sync()) { + if (chunk_is_full_1 == 0) { + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + mbarrier_arrive(v_full_addr + (load_stage) * 8); + } + } + load_stage += 1; + if (load_stage == 5) { load_stage = 0; _phase_v_free ^= 1; _phase_qk_full_2 ^= 1; } + } + } + // ---- Role: persistent owner mailbox ingress coordinator ---- + } else if (cta_rank < kOwnerCount && warp == 12) { + int task_idx_4 = bid; + int seq_idx_4 = seq_order[task_idx_4 / num_heads]; + int seq_len_4 = int(cu_seqlens[seq_idx_4 + 1] - cu_seqlens[seq_idx_4]); + int num_chunks_4 = (seq_len_4 + 31) / 32; + unsigned int mailbox_stage = 0; + unsigned int raw_phase = 1; + unsigned int smem_phase = 1; + for (int chunk_idx_3 = 0; chunk_idx_3 < num_chunks_4; + ++chunk_idx_3) { + mbarrier_wait( + raw_inputs_free_addr + mailbox_stage * 8, raw_phase); + mbarrier_wait( + smem_free_addr + mailbox_stage * 8, smem_phase); + long long packet_idx = + (long long)task_idx_4 * mailbox_depth + + chunk_idx_3 % mailbox_depth; + const unsigned int generation = + (unsigned int)(chunk_idx_3 / mailbox_depth); + wait_k1_global_ready( + k1_flags + packet_idx, (generation << 3) | 1u); + if (elect_sync()) { + load_k1_from_global( + smem_qd_addr + mailbox_stage * 41984, + qk_full_addr + mailbox_stage * 8, + k1_workspace + packet_idx * kK1PacketBytes); + } + mailbox_stage += 1; + if (mailbox_stage == 5) { + mailbox_stage = 0; + raw_phase ^= 1; + smem_phase ^= 1; + } + } + // ---- Role: prep ---- + } else if (cta_rank >= kProducerFirstRank && warp >= 12 && warp <= 31) { + asm volatile("setmaxnreg.dec.sync.aligned.u32 48;"); + { // prep_main + int task_idx_4 = bid; + int seq_idx_4 = seq_order[task_idx_4 / num_heads]; + int head_idx_3 = task_idx_4 % num_heads; + long long bos_4 = cu_seqlens[seq_idx_4]; + long long eos_4 = cu_seqlens[seq_idx_4 + 1]; + int seq_len_4 = (int)(eos_4 - bos_4); + int num_chunks_4 = (seq_len_4 + 32 - 1) / 32; + int instance_id = (warp - 12) / 4; + int prep_instance = instance_id; + int warp_id_in_role_2 = (warp - 12); + int prep_local_warp = warp_id_in_role_2 - prep_instance * 4; + int prep_tid = prep_local_warp * 32 + lane; + const int first_work = + (cta_rank - kProducerFirstRank) * 5 + prep_instance; + const int producer_stride = kProducerCount * 5; + const int total_work = num_chunks_4; + int num_prep_iters = first_work < total_work + ? (total_work - 1 - first_work) / producer_stride + 1 + : 0; + unsigned int prep_stage = (unsigned int)prep_instance; + int gate_rate_stage_f32 = prep_instance * 10496; + if (prep_tid == 0) { + float _expf_0 = __expf(A_log[head_idx_3]); + smem_gate_rate_all[gate_rate_stage_f32] = _expf_0; + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + unsigned int _phase_raw_inputs_free = 1; + unsigned int _phase_gate_raw_full = 0; + unsigned int _phase_smem_free = 1; + unsigned int _phase_qk_raw_full = 0; + unsigned int _phase_prep_diag_ready = 0; + unsigned int _phase_prep_inv16_ready = 0; + #pragma unroll 1 + for (int prep_iter = 0; prep_iter < num_prep_iters; prep_iter++) { + const int work_idx = prep_iter * producer_stride + first_work; + int chunk_idx_3 = work_idx; + int stage_f32 = prep_stage * 10496; + int stage_bf16 = prep_stage * 20992; + int chunk_is_full_2 = ((seq_len_4 >= (chunk_idx_3 + 1) * 32) ? 1 : 0); + float early_beta_value = 0.0f; + float early_gate0 = 0.0f; + if (chunk_is_full_2 != 0) { + mbarrier_wait(raw_inputs_free_addr + (prep_stage) * 8, _phase_raw_inputs_free); + if (prep_local_warp == 0) { + if (elect_sync()) { + mbarrier_arrive_expect_tx(gate_raw_full_addr + (prep_stage) * 8, 8704); + tma_3d_gmem2smem(smem_g_raw_addr + prep_stage * 41984, g_tma, 0, head_idx_3, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), gate_raw_full_addr + (prep_stage) * 8); + tma_2d_gmem2smem(smem_beta_raw_addr + prep_stage * 41984, beta_tma, head_idx_3 / 8 * 8, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), gate_raw_full_addr + (prep_stage) * 8); + mbarrier_arrive_expect_tx(qk_raw_full_addr + (prep_stage) * 8, 16384); + tma_4d_gmem2smem(smem_kd_addr + prep_stage * 41984, k_tma, 0, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), head_idx_3, 0, qk_raw_full_addr + (prep_stage) * 8); + } + } + mbarrier_wait(gate_raw_full_addr + (prep_stage) * 8, _phase_gate_raw_full); + if (prep_local_warp == 2 && lane < 32) { + unsigned int beta_raw_pair[1]; + asm volatile("ld.shared.b32 %0, [%1];" : "=r"(*reinterpret_cast(&beta_raw_pair[0])) : "r"(smem_beta_raw_addr + prep_stage * 41984 + (unsigned int)(lane * 16) + (unsigned int)(head_idx_3 % 8 / 2 * 4))); + float beta_raw_pair_fp32[2]; + #pragma unroll + for (int _pair = 0; _pair < 1; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&beta_raw_pair_fp32[_pair * 2])[0]), "=f"((&beta_raw_pair_fp32[_pair * 2])[1]) + : "r"(beta_raw_pair[_pair + 0])); + } + float beta_logit = beta_raw_pair_fp32[0]; + if (head_idx_3 % 2 != 0) { + beta_logit = beta_raw_pair_fp32[1]; + } + float _tanh_approx_0; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_0) : "f"(beta_logit * 0.5f)); + early_beta_value = _tanh_approx_0 * 0.5f + 0.5f; + } + if (prep_tid < 128) { + float early_gate_rate = smem_gate_rate_all[stage_f32]; + float early_gate_bias = dt_bias[head_idx_3 * 128 + prep_tid]; + __nv_bfloat16 early_gate_raw = smem_g_raw_all[stage_bf16 + prep_tid]; + float _cvt_f32_0 = __bfloat162float(early_gate_raw); + float early_gate_arg = early_gate_rate * (_cvt_f32_0 + early_gate_bias); + float _tanh_approx_1; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_1) : "f"(early_gate_arg * 0.5f)); + float early_gate_sigmoid = _tanh_approx_1 * 0.5f + 0.5f; + early_gate0 = lower_bound * 1.4426950408889634f * early_gate_sigmoid; + } + } + mbarrier_wait(smem_free_addr + (prep_stage) * 8, _phase_smem_free); + if (chunk_is_full_2 != 0) { + if (prep_local_warp == 0) { + if (elect_sync()) { + tma_4d_gmem2smem(smem_q_raw_prefetch_addr + prep_stage * 41984, q_tma, 0, (int)(bos_4 + (long long)(chunk_idx_3 * 32)), head_idx_3, 0, qk_raw_full_addr + (prep_stage) * 8); + } + } + } + if (chunk_is_full_2 == 0) { + #pragma unroll + for (int gate_load_pass = 0; gate_load_pass < 4; gate_load_pass++) { + int gate_load_item = gate_load_pass * 128 + prep_tid; + int gate_load_row = gate_load_item / 16; + int gate_load_segment = gate_load_item % 16; + long long gate_load_token = bos_4 + (long long)(chunk_idx_3 * 32 + gate_load_row); + long long gate_load_base = (gate_load_token * (long long)num_heads + (long long)head_idx_3) * 128 + (long long)(gate_load_segment * 8); + asm volatile("cp.async.cg.shared::cta.global [%0], [%1], 16, %2;" + :: "r"(smem_g_raw_addr + prep_stage * 41984 + (unsigned int)(gate_load_item * 16)), "l"(g + gate_load_base), "r"((gate_load_token < eos_4) ? 16 : 0)); + } + } + if (chunk_is_full_2 == 0) { + asm volatile("cp.async.commit_group;"); + asm volatile("cp.async.wait_group 0;"); + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + } + if (prep_local_warp == 2 && lane < 32) { + float beta_value = early_beta_value; + if (chunk_is_full_2 == 0) { + long long beta_token = bos_4 + (long long)(chunk_idx_3 * 32 + lane); + if (beta_token < eos_4) { + float beta_logit_1 = (float)beta[beta_token * (long long)num_heads + (long long)head_idx_3]; + float _tanh_approx_2; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_2) : "f"(beta_logit_1 * 0.5f)); + beta_value = _tanh_approx_2 * 0.5f + 0.5f; + } + } + smem_prep_beta_all[stage_f32 + lane] = beta_value; + } + if (prep_tid < 128) { + int gate_col = prep_tid; + float gate_rate = smem_gate_rate_all[stage_f32]; + float gate_bias = dt_bias[head_idx_3 * 128 + gate_col]; + float prefix_log2 = 0.0f; + for (int gate_row = 0; gate_row < 32; gate_row++) { + long long gate_token = bos_4 + (long long)(chunk_idx_3 * 32 + gate_row); + float gate_log2 = 0.0f; + int gate_needs_compute = 1; + if (gate_row == 0) { + if (chunk_is_full_2 != 0) { + gate_log2 = early_gate0; + gate_needs_compute = 0; + } + } + if (gate_needs_compute != 0) { + if (gate_token < eos_4) { + __nv_bfloat16 gate_raw = smem_g_raw_all[stage_bf16 + gate_row * 128 + gate_col]; + float _cvt_f32_1 = __bfloat162float(gate_raw); + float gate_arg = gate_rate * (_cvt_f32_1 + gate_bias); + float _tanh_approx_3; + asm volatile("tanh.approx.f32 %0, %1;" : "=f"(_tanh_approx_3) : "f"(gate_arg * 0.5f)); + float gate_sigmoid = _tanh_approx_3 * 0.5f + 0.5f; + gate_log2 = lower_bound * 1.4426950408889634f * gate_sigmoid; + } + } + prefix_log2 += gate_log2; + smem_gate_all[stage_f32 + gate_row * 128 + gate_col] = prefix_log2; + } + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + if (chunk_is_full_2 != 0) { + mbarrier_wait(qk_raw_full_addr + (prep_stage) * 8, _phase_qk_raw_full); + } + if (prep_tid < 128) { + float total_log2 = smem_gt_prefix_all[stage_f32 + prep_tid]; + float _exp2_0 = approx_exp2(total_log2 - lower_bound * 1.4426950408889634f * 16.0f); + smem_restore_factor_all[stage_f32 + prep_tid] = _exp2_0; + } + if (prep_tid == 0) { + float _exp2_1 = approx_exp2(lower_bound * 1.4426950408889634f * 16.0f); + smem_restore_factor_all[stage_f32 + 128] = _exp2_1; + } + #pragma unroll 1 + for (int work_pass = 0; work_pass < 4; work_pass++) { + int work_item = work_pass * 128 + prep_tid; + int row_1 = work_item / 16; + int segment_1 = work_item % 16; + long long token_1 = bos_4 + (long long)(chunk_idx_3 * 32 + row_1); + int token_valid_1 = ((token_1 < eos_4) ? 1 : 0); + long long gmem_base = (token_1 * (long long)num_heads + (long long)head_idx_3) * 128 + (long long)(segment_1 * 8); + float q_raw_vec[8]; + float k_raw_vec[8]; + q_raw_vec[0] = 0.0f; + q_raw_vec[1] = 0.0f; + q_raw_vec[2] = 0.0f; + q_raw_vec[3] = 0.0f; + q_raw_vec[4] = 0.0f; + q_raw_vec[5] = 0.0f; + q_raw_vec[6] = 0.0f; + q_raw_vec[7] = 0.0f; + k_raw_vec[0] = 0.0f; + k_raw_vec[1] = 0.0f; + k_raw_vec[2] = 0.0f; + k_raw_vec[3] = 0.0f; + k_raw_vec[4] = 0.0f; + k_raw_vec[5] = 0.0f; + k_raw_vec[6] = 0.0f; + k_raw_vec[7] = 0.0f; + if (chunk_is_full_2 != 0) { + unsigned int packed[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed[0])), "=r"(*reinterpret_cast(&packed[(0) + 1])), "=r"(*reinterpret_cast(&packed[(0) + 2])), "=r"(*reinterpret_cast(&packed[(0) + 3])) + : "r"((smem_q_raw_prefetch_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32[_pair * 2])[0]), "=f"((&packed_fp32[_pair * 2])[1]) + : "r"(packed[_pair + 0])); + } + #pragma unroll + for (int value_idx = 0; value_idx < 8; value_idx++) { + q_raw_vec[value_idx] = packed_fp32[value_idx]; + } + unsigned int packed_0[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0[0])), "=r"(*reinterpret_cast(&packed_0[(0) + 1])), "=r"(*reinterpret_cast(&packed_0[(0) + 2])), "=r"(*reinterpret_cast(&packed_0[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32[_pair * 2])[0]), "=f"((&packed_0_fp32[_pair * 2])[1]) + : "r"(packed_0[_pair + 0])); + } + #pragma unroll + for (int value_idx_1 = 0; value_idx_1 < 8; value_idx_1++) { + k_raw_vec[value_idx_1] = packed_0_fp32[value_idx_1]; + } + } else if (token_valid_1 != 0) { + { + const uint4* _vptr_0 = reinterpret_cast(q + gmem_base); + uint4 _vld_0[1]; + #pragma unroll + for (int _blk = 0; _blk < 1; _blk++) { + _vld_0[_blk] = _vptr_0[_blk]; + uint32_t* _vpairs_0 = reinterpret_cast(&_vld_0[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&q_raw_vec[0 + _blk * 8 + _pair * 2])[0]), "=f"((&q_raw_vec[0 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_0[_pair])); + } + } + } + { + const uint4* _vptr_1 = reinterpret_cast(k + gmem_base); + uint4 _vld_1[1]; + #pragma unroll + for (int _blk = 0; _blk < 1; _blk++) { + _vld_1[_blk] = _vptr_1[_blk]; + uint32_t* _vpairs_1 = reinterpret_cast(&_vld_1[_blk]); + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&k_raw_vec[0 + _blk * 8 + _pair * 2])[0]), "=f"((&k_raw_vec[0 + _blk * 8 + _pair * 2])[1]) + : "r"(_vpairs_1[_pair])); + } + } + } + } + float q_sum = 0.0f; + float k_sum = 0.0f; + for (int elem_in_segment = 0; elem_in_segment < 8; elem_in_segment++) { + float q_raw = q_raw_vec[elem_in_segment]; + float k_raw = k_raw_vec[elem_in_segment]; + float _fma_0 = __fmaf_rn(q_raw, q_raw, q_sum); + q_sum = _fma_0; + float _fma_1 = __fmaf_rn(k_raw, k_raw, k_sum); + k_sum = _fma_1; + } + float _shfl_xor_0 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 8); + q_sum += _shfl_xor_0; + float _shfl_xor_1 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 8); + k_sum += _shfl_xor_1; + float _shfl_xor_2 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 4); + q_sum += _shfl_xor_2; + float _shfl_xor_3 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 4); + k_sum += _shfl_xor_3; + float _shfl_xor_4 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 2); + q_sum += _shfl_xor_4; + float _shfl_xor_5 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 2); + k_sum += _shfl_xor_5; + float _shfl_xor_6 = __shfl_xor_sync(0xFFFFFFFF, q_sum, 1); + q_sum += _shfl_xor_6; + float _shfl_xor_7 = __shfl_xor_sync(0xFFFFFFFF, k_sum, 1); + k_sum += _shfl_xor_7; + float _rsqrt_0 = rsqrtf(q_sum + 1e-06f); + float q_inv = _rsqrt_0; + float _rsqrt_1 = rsqrtf(k_sum + 1e-06f); + float k_inv = _rsqrt_1; + const float2 _scale2_2 = {q_inv, q_inv}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(q_raw_vec)[_ls], _scale2_2); + const float2 _scale2_3 = {k_inv, k_inv}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(k_raw_vec)[_ls], _scale2_3); + float qd_vec[8]; + float kd_vec[8]; + float ki_vec[8]; + for (int elem_in_segment_1 = 0; elem_in_segment_1 < 8; elem_in_segment_1++) { + int col = segment_1 * 8 + elem_in_segment_1; + float prefix = smem_gate_all[stage_f32 + row_1 * 128 + col]; + float common_log2 = lower_bound * 1.4426950408889634f * 16.0f; + float _exp2_2 = approx_exp2(prefix - common_log2); + float decay = _exp2_2; + qd_vec[elem_in_segment_1] = decay; + kd_vec[elem_in_segment_1] = decay; + ki_vec[elem_in_segment_1] = k_raw_vec[elem_in_segment_1] / decay; + } + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(qd_vec)[_ls], reinterpret_cast(q_raw_vec)[_ls]); + const float2 _scale2_4 = {scale, scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(qd_vec)[_ls], _scale2_4); + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(kd_vec)[_ls], reinterpret_cast(k_raw_vec)[_ls]); + unsigned int packed_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(qd_vec[_lp*2 + 0], qd_vec[_lp*2+1 + 0])); + packed_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word = 0; word < 4; word++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word * 4)), "r"(packed_1[word])); + } + unsigned int packed_0_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(kd_vec[_lp*2 + 0], kd_vec[_lp*2+1 + 0])); + packed_0_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_1 = 0; word_1 < 4; word_1++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_1 * 4)), "r"(packed_0_1[word_1])); + } + unsigned int packed_1_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(ki_vec[_lp*2 + 0], ki_vec[_lp*2+1 + 0])); + packed_1_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_2 = 0; word_2 < 4; word_2++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 ^ (segment_1 * 8 / 64 * 4096 + row_1 * 128 + segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_2 * 4)), "r"(packed_1_1[word_2])); + } + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + int pair_row_base = prep_local_warp / 2 * 16; + int pair_col_base = prep_local_warp % 2 * 16; + unsigned int a_frag[4]; + unsigned int b_frag[4]; + float acc[8]; + if (pair_row_base >= pair_col_base) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[0]), "=f"(acc[1]), "=f"(acc[2]), "=f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[4]), "=f"(acc[(4) + 1]), "=f"(acc[(4) + 2]), "=f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_kd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + int row0 = pair_row_base + lane / 4; + int row1 = row0 + 8; + int col0 = pair_col_base + lane % 4 * 2; + float beta0 = smem_prep_beta_all[stage_f32 + row0]; + float beta1 = smem_prep_beta_all[stage_f32 + row1]; + float seed[8]; + seed[0] = 0.0f; + seed[1] = 0.0f; + seed[2] = 0.0f; + seed[3] = 0.0f; + seed[4] = 0.0f; + seed[5] = 0.0f; + seed[6] = 0.0f; + seed[7] = 0.0f; + if (row0 > col0) { + seed[0] = acc[0] * beta0; + } + if (row0 > col0 + 1) { + seed[1] = acc[1] * beta0; + } + if (row1 > col0) { + seed[2] = acc[2] * beta1; + } + if (row1 > col0 + 1) { + seed[3] = acc[3] * beta1; + } + if (row0 > col0 + 8) { + seed[4] = acc[4] * beta0; + } + if (row0 > col0 + 9) { + seed[5] = acc[5] * beta0; + } + if (row1 > col0 + 8) { + seed[6] = acc[6] * beta1; + } + if (row1 > col0 + 9) { + seed[7] = acc[7] * beta1; + } + unsigned int seed_packed[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(seed[_lp*2 + 0], seed[_lp*2+1 + 0])); + seed_packed[_lp] = *(uint32_t*)&_bf2; + } + int seed_lane_row = lane % 16; + int seed_lane_col = lane / 16 * 8; + int byte_off = (pair_row_base + seed_lane_row) * 128 + (pair_col_base + seed_lane_col) * 2; + int swizzled_off = byte_off ^ (byte_off >> 7 & 7) << 4; + int seed_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off; + uint32_t _stmatrix_addr_5 = static_cast((unsigned long long)seed_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_5), "r"(*reinterpret_cast(&seed_packed[0])), "r"(*reinterpret_cast(&seed_packed[1])), "r"(*reinterpret_cast(&seed_packed[2])), "r"(*reinterpret_cast(&seed_packed[3])) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16) * 16)) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[0]), "=f"(acc[1]), "=f"(acc[2]), "=f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(acc[4]), "=f"(acc[(4) + 1]), "=f"(acc[(4) + 2]), "=f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a_frag[0]), "=r"(a_frag[1]), "=r"(a_frag[2]), "=r"(a_frag[3]) + : "r"(smem_qd_addr + prep_stage * 41984 + (unsigned int)(((lane / 16 / 8 * 256 + (pair_row_base + lane % 16) * 8 + (lane / 16 % 8 * 16 ^ (pair_row_base + lane % 16 & 7) << 4) / 16 ^ 2 ^ 6 ^ 2 ^ 6) + 256 ^ 2 ^ 6 ^ 2) * 16)) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(b_frag[0]), "=r"(b_frag[1]), "=r"(b_frag[2]), "=r"(b_frag[3]) + : "r"(smem_ki_addr + prep_stage * 41984 + (unsigned int)(((((((((lane % 16 / 8 / 8 * 256 + (pair_col_base + 8 * (lane / 16) + lane % 8) * 8 + (lane % 16 / 8 % 8 * 16 ^ (pair_col_base + 8 * (lane / 16) + lane % 8 & 7) << 4) / 16 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256 + 256 ^ 6) + 256 - 256 + 256 ^ 2) - 256 + 256 ^ 6) - 256 + 256 ^ 2) - 256) * 16)) + : "memory"); + 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"(acc[0]), "+f"(acc[1]), "+f"(acc[2]), "+f"(acc[3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[0]), "r"(b_frag[1])); + 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"(acc[4]), "+f"(acc[(4) + 1]), "+f"(acc[(4) + 2]), "+f"(acc[(4) + 3]) + : "r"(a_frag[0]), "r"(a_frag[1]), "r"(a_frag[2]), "r"(a_frag[3]), "r"(b_frag[2]), "r"(b_frag[(2) + 1])); + } else { + acc[0] = 0.0f; + acc[1] = 0.0f; + acc[2] = 0.0f; + acc[3] = 0.0f; + acc[4] = 0.0f; + acc[5] = 0.0f; + acc[6] = 0.0f; + acc[7] = 0.0f; + } + int row0_1 = pair_row_base + lane / 4; + int row1_1 = row0_1 + 8; + int col0_1 = pair_col_base + lane % 4 * 2; + float mqk[8]; + mqk[0] = 0.0f; + mqk[1] = 0.0f; + mqk[2] = 0.0f; + mqk[3] = 0.0f; + mqk[4] = 0.0f; + mqk[5] = 0.0f; + mqk[6] = 0.0f; + mqk[7] = 0.0f; + if (row0_1 >= col0_1) { + mqk[0] = acc[0]; + } + if (row0_1 >= col0_1 + 1) { + mqk[1] = acc[1]; + } + if (row1_1 >= col0_1) { + mqk[2] = acc[2]; + } + if (row1_1 >= col0_1 + 1) { + mqk[3] = acc[3]; + } + if (row0_1 >= col0_1 + 8) { + mqk[4] = acc[4]; + } + if (row0_1 >= col0_1 + 9) { + mqk[5] = acc[5]; + } + if (row1_1 >= col0_1 + 8) { + mqk[6] = acc[6]; + } + if (row1_1 >= col0_1 + 9) { + mqk[7] = acc[7]; + } + unsigned int mqk_packed[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(mqk[_lp*2 + 0], mqk[_lp*2+1 + 0])); + mqk_packed[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int publish_pair = 0; publish_pair < 2; publish_pair++) { + int publish_row = pair_col_base + publish_pair * 8 + (lane & 7); + int publish_col = 128 + pair_row_base + lane / 8 * 8; + uint32_t _stmatrix_addr_6 = static_cast((unsigned long long)(smem_final_trans_addr + prep_stage * 41984 + (unsigned int)(publish_col / 64 * 4096 + publish_row * 128 + publish_col % 64 * 2 ^ (publish_col / 64 * 4096 + publish_row * 128 + publish_col % 64 * 2 >> 7 & 7) << 4))); + asm volatile("stmatrix.sync.aligned.m8n8.x2.trans.shared.b16 [%0], {%1, %2};\n" + :: "r"(_stmatrix_addr_6), "r"(*reinterpret_cast(&mqk_packed[publish_pair * 2])), "r"(*reinterpret_cast(&mqk_packed[publish_pair * 2 + 1])) + : "memory"); + } + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + if (prep_tid < 128) { + float total_log2_1 = smem_gt_prefix_all[stage_f32 + prep_tid]; + float _exp2_3 = approx_exp2(total_log2_1); + smem_gt_all[stage_f32 + prep_tid] = _exp2_3; + } + if (prep_local_warp >= 2) { + int stage_f32_0 = prep_stage * 10496; + float restore_scale = smem_restore_factor_all[stage_f32_0 + 128]; + float restore_factor[8]; + int restore_segment = lane & 15; + #pragma unroll + for (int restore_elem = 0; restore_elem < 8; restore_elem++) { + int restore_col = restore_segment * 8 + restore_elem; + restore_factor[restore_elem] = smem_restore_factor_all[stage_f32_0 + restore_col]; + } + #pragma unroll 1 + for (int restore_pass = 0; restore_pass < 6; restore_pass++) { + int restore_row = 8 + (prep_local_warp - 2) * 12 + restore_pass * 2 + (lane >> 4); + float restore_qd_values[8]; + float restore_kd_values[8]; + float restore_ki_values[8]; + unsigned int packed_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_2[0])), "=r"(*reinterpret_cast(&packed_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_2[(0) + 3])) + : "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_1[_pair * 2])[0]), "=f"((&packed_fp32_1[_pair * 2])[1]) + : "r"(packed_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_2 = 0; value_idx_2 < 8; value_idx_2++) { + restore_qd_values[value_idx_2] = packed_fp32_1[value_idx_2]; + } + unsigned int packed_0_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0_2[0])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_0_2[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32_1[_pair * 2])[0]), "=f"((&packed_0_fp32_1[_pair * 2])[1]) + : "r"(packed_0_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_3 = 0; value_idx_3 < 8; value_idx_3++) { + restore_kd_values[value_idx_3] = packed_0_fp32_1[value_idx_3]; + } + unsigned int packed_1_2[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_1_2[0])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 1])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 2])), "=r"(*reinterpret_cast(&packed_1_2[(0) + 3])) + : "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_1_fp32[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_1_fp32[_pair * 2])[0]), "=f"((&packed_1_fp32[_pair * 2])[1]) + : "r"(packed_1_2[_pair + 0])); + } + #pragma unroll + for (int value_idx_4 = 0; value_idx_4 < 8; value_idx_4++) { + restore_ki_values[value_idx_4] = packed_1_fp32[value_idx_4]; + } + float restore_kr_values[8]; + #pragma unroll + for (int restore_elem_1 = 0; restore_elem_1 < 8; restore_elem_1++) { + restore_kr_values[restore_elem_1] = restore_ki_values[restore_elem_1] * restore_factor[restore_elem_1]; + } + const float2 _scale2_7 = {restore_scale, restore_scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_qd_values)[_ls], _scale2_7); + const float2 _scale2_8 = {restore_scale, restore_scale}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_kd_values)[_ls], _scale2_8); + unsigned int packed_2_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_qd_values[_lp*2 + 0], restore_qd_values[_lp*2+1 + 0])); + packed_2_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_3 = 0; word_3 < 4; word_3++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_3 * 4)), "r"(packed_2_1[word_3])); + } + unsigned int packed_3[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kd_values[_lp*2 + 0], restore_kd_values[_lp*2+1 + 0])); + packed_3[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_4 = 0; word_4 < 4; word_4++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_4 * 4)), "r"(packed_3[word_4])); + } + unsigned int packed_4[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kr_values[_lp*2 + 0], restore_kr_values[_lp*2+1 + 0])); + packed_4[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_5 = 0; word_5 < 4; word_5++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kr_trans_addr + prep_stage * 41984 + (unsigned int)(restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 ^ (restore_segment * 8 / 64 * 4096 + restore_row * 128 + restore_segment * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_5 * 4)), "r"(packed_4[word_5])); + } + } + } + if (prep_local_warp == 0) { + int inverse_row = lane; + int diag_block = inverse_row / 8; + int lane_in_diag = lane & 7; + float inv_row[8]; + unsigned int packed_5[4]; + int byte_off_1 = inverse_row * 128 + diag_block * 8 * 2; + int swizzled_off_1 = byte_off_1 ^ (byte_off_1 >> 7 & 7) << 4; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_5[0])), "=r"(*reinterpret_cast(&packed_5[(0) + 1])), "=r"(*reinterpret_cast(&packed_5[(0) + 2])), "=r"(*reinterpret_cast(&packed_5[(0) + 3])) + : "r"(smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1)); + float packed_fp32_2[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_2[_pair * 2])[0]), "=f"((&packed_fp32_2[_pair * 2])[1]) + : "r"(packed_5[_pair + 0])); + } + #pragma unroll + for (int value_idx_5 = 0; value_idx_5 < 8; value_idx_5++) { + inv_row[value_idx_5] = packed_fp32_2[value_idx_5]; + } + #pragma unroll + for (int diag_elem = 0; diag_elem < 8; diag_elem++) { + if (lane_in_diag == diag_elem) { + inv_row[diag_elem] = 1.0f; + } + } + int diag_group_base = lane - lane_in_diag; + #pragma unroll + for (int src_row = 0; src_row < 7; src_row++) { + float row_scale = -inv_row[src_row]; + #pragma unroll + for (int prev_col = 0; prev_col < src_row; prev_col++) { + int pivot_lane = diag_group_base + src_row; + float _shfl_0 = __shfl_sync(0xFFFFFFFF, inv_row[prev_col], pivot_lane); + float pivot = _shfl_0; + if (lane_in_diag > src_row) { + float _fma_2 = __fmaf_rn(row_scale, pivot, inv_row[prev_col]); + inv_row[prev_col] = _fma_2; + } + } + if (lane_in_diag > src_row) { + inv_row[src_row] = row_scale; + } + } + unsigned int packed_0_3[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(inv_row[_lp*2 + 0], inv_row[_lp*2+1 + 0])); + packed_0_3[_lp] = *(uint32_t*)&_bf2; + } + int byte_off_1_1 = inverse_row * 128 + diag_block * 8 * 2; + int swizzled_off_2 = byte_off_1_1 ^ (byte_off_1_1 >> 7 & 7) << 4; + #pragma unroll + for (int word_6 = 0; word_6 < 4; word_6++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"(smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_2 + (unsigned int)(word_6 * 4)), "r"(packed_0_3[word_6])); + } + } + if (prep_local_warp < 2) { + if (elect_sync()) { + mbarrier_arrive(prep_diag_ready_addr + (prep_stage) * 8); + } + mbarrier_wait(prep_diag_ready_addr + (prep_stage) * 8, _phase_prep_diag_ready); + } + if (prep_local_warp < 2) { + int lane_row = lane & 7; + int byte_off_2 = (prep_local_warp * 16 + 8 + lane_row) * 128 + (prep_local_warp * 16 + 8) * 2; + int swizzled_off_3 = byte_off_2 ^ (byte_off_2 >> 7 & 7) << 4; + int d_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3; + int byte_off_0 = (prep_local_warp * 16 + 8 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_1_1 = byte_off_0 ^ (byte_off_0 >> 7 & 7) << 4; + int c_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1_1; + int byte_off_2_1 = (prep_local_warp * 16 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_3_1 = byte_off_2_1 ^ (byte_off_2_1 >> 7 & 7) << 4; + int a_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3_1; + unsigned int d_frag[2]; + unsigned int c_frag[1]; + float dc_acc[4]; + unsigned int dc_bf16[2]; + unsigned int inv_a_frag[1]; + float o_acc[4]; + unsigned int o_bf16[2]; + asm volatile("ldmatrix.sync.aligned.m8n8.x1.shared.b16 {%0}, [%1];\n" + : "=r"(d_frag[0]) + : "r"(d_addr) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x1.shared.b16 {%0}, [%1];\n" + : "=r"(d_frag[1]) + : "r"(d_addr) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x1.trans.shared.b16 {%0}, [%1];\n" + : "=r"(c_frag[0]) + : "r"(c_addr) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k8.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5}, {%6}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc_acc[0]), "=f"(dc_acc[1]), "=f"(dc_acc[2]), "=f"(dc_acc[3]) + : "r"(d_frag[0]), "r"(d_frag[1]), "r"(c_frag[0])); + const float2 _scale2_9 = {-1.0f, -1.0f}; + #pragma unroll + for (int _ls = 0; _ls < 2; _ls++) + mul_f32x2_inplace(&reinterpret_cast(dc_acc)[_ls], _scale2_9); + #pragma unroll + for (int _lp = 0; _lp < 2; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(dc_acc[_lp*2 + 0], dc_acc[_lp*2+1 + 0])); + dc_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile("ldmatrix.sync.aligned.m8n8.x1.trans.shared.b16 {%0}, [%1];\n" + : "=r"(inv_a_frag[0]) + : "r"(a_addr) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k8.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5}, {%6}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o_acc[0]), "=f"(o_acc[1]), "=f"(o_acc[2]), "=f"(o_acc[3]) + : "r"(dc_bf16[0]), "r"(dc_bf16[1]), "r"(inv_a_frag[0])); + #pragma unroll + for (int _lp = 0; _lp < 2; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(o_acc[_lp*2 + 0], o_acc[_lp*2+1 + 0])); + o_bf16[_lp] = *(uint32_t*)&_bf2; + } + int byte_off_4 = (prep_local_warp * 16 + 8 + lane_row) * 128 + prep_local_warp * 16 * 2; + int swizzled_off_5 = byte_off_4 ^ (byte_off_4 >> 7 & 7) << 4; + int o_addr = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_5; + uint32_t _stmatrix_addr_10 = static_cast((unsigned long long)o_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x1.shared.b16 [%0], {%1};\n" + :: "r"(_stmatrix_addr_10), "r"(*reinterpret_cast(&o_bf16[0])) + : "memory"); + if (elect_sync()) { + mbarrier_arrive(prep_inv16_ready_addr + (prep_stage) * 8); + } + mbarrier_wait(prep_inv16_ready_addr + (prep_stage) * 8, _phase_prep_inv16_ready); + } + if (prep_local_warp == 0) { + int lane_row_1 = lane % 16; + int lane_col = lane / 16 * 8; + int byte_off_3 = (16 + lane_row_1) * 128 + (16 + lane_col) * 2; + int swizzled_off_4 = byte_off_3 ^ (byte_off_3 >> 7 & 7) << 4; + int d_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_4; + int byte_off_0_1 = (16 + lane_row_1) * 128 + lane_col * 2; + int swizzled_off_1_2 = byte_off_0_1 ^ (byte_off_0_1 >> 7 & 7) << 4; + int c_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_1_2; + int byte_off_2_2 = lane_row_1 * 128 + lane_col * 2; + int swizzled_off_3_2 = byte_off_2_2 ^ (byte_off_2_2 >> 7 & 7) << 4; + int a_addr_1 = smem_inv_work_addr + prep_stage * 41984 + (unsigned int)swizzled_off_3_2; + unsigned int d32_frag[4]; + unsigned int c32_frag[4]; + float dc32_acc[8]; + unsigned int dc32_bf16[4]; + unsigned int a32_frag[4]; + float o32_acc[8]; + unsigned int o32_bf16[4]; + unsigned int zero32_bf16[4]; + asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(d32_frag[0]), "=r"(d32_frag[1]), "=r"(d32_frag[2]), "=r"(d32_frag[3]) + : "r"(d_addr_1) + : "memory"); + int d_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)((16 + lane_col) / 16 * 1024 + (16 + lane_row_1) * 32 + (16 + lane_col) % 16 * 2 ^ ((16 + lane_col) / 16 * 1024 + (16 + lane_row_1) * 32 + (16 + lane_col) % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_11 = static_cast((unsigned long long)d_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_11), "r"(*reinterpret_cast(&d32_frag[0])), "r"(*reinterpret_cast(&d32_frag[1])), "r"(*reinterpret_cast(&d32_frag[2])), "r"(*reinterpret_cast(&d32_frag[3])) + : "memory"); + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(c32_frag[0]), "=r"(c32_frag[1]), "=r"(c32_frag[2]), "=r"(c32_frag[3]) + : "r"(c_addr_1) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc32_acc[0]), "=f"(dc32_acc[1]), "=f"(dc32_acc[2]), "=f"(dc32_acc[3]) + : "r"(d32_frag[0]), "r"(d32_frag[1]), "r"(d32_frag[2]), "r"(d32_frag[3]), "r"(c32_frag[0]), "r"(c32_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(dc32_acc[4]), "=f"(dc32_acc[(4) + 1]), "=f"(dc32_acc[(4) + 2]), "=f"(dc32_acc[(4) + 3]) + : "r"(d32_frag[0]), "r"(d32_frag[1]), "r"(d32_frag[2]), "r"(d32_frag[3]), "r"(c32_frag[2]), "r"(c32_frag[(2) + 1])); + const float2 _scale2_12 = {-1.0f, -1.0f}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(dc32_acc)[_ls], _scale2_12); + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(dc32_acc[_lp*2 + 0], dc32_acc[_lp*2+1 + 0])); + dc32_bf16[_lp] = *(uint32_t*)&_bf2; + } + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16 {%0, %1, %2, %3}, [%4];\n" + : "=r"(a32_frag[0]), "=r"(a32_frag[1]), "=r"(a32_frag[2]), "=r"(a32_frag[3]) + : "r"(a_addr_1) + : "memory"); + int a_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)(lane_col / 16 * 1024 + lane_row_1 * 32 + lane_col % 16 * 2 ^ (lane_col / 16 * 1024 + lane_row_1 * 32 + lane_col % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_13 = static_cast((unsigned long long)a_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.trans.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_13), "r"(*reinterpret_cast(&a32_frag[0])), "r"(*reinterpret_cast(&a32_frag[1])), "r"(*reinterpret_cast(&a32_frag[2])), "r"(*reinterpret_cast(&a32_frag[3])) + : "memory"); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o32_acc[0]), "=f"(o32_acc[1]), "=f"(o32_acc[2]), "=f"(o32_acc[3]) + : "r"(dc32_bf16[0]), "r"(dc32_bf16[1]), "r"(dc32_bf16[2]), "r"(dc32_bf16[3]), "r"(a32_frag[0]), "r"(a32_frag[1])); + asm volatile("mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 {%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {0f00000000, 0f00000000, 0f00000000, 0f00000000};\n" + : "=f"(o32_acc[4]), "=f"(o32_acc[(4) + 1]), "=f"(o32_acc[(4) + 2]), "=f"(o32_acc[(4) + 3]) + : "r"(dc32_bf16[0]), "r"(dc32_bf16[1]), "r"(dc32_bf16[2]), "r"(dc32_bf16[3]), "r"(a32_frag[2]), "r"(a32_frag[(2) + 1])); + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(o32_acc[_lp*2 + 0], o32_acc[_lp*2+1 + 0])); + o32_bf16[_lp] = *(uint32_t*)&_bf2; + } + int o_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)(lane_col / 16 * 1024 + (16 + lane_row_1) * 32 + lane_col % 16 * 2 ^ (lane_col / 16 * 1024 + (16 + lane_row_1) * 32 + lane_col % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_14 = static_cast((unsigned long long)o_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_14), "r"(*reinterpret_cast(&o32_bf16[0])), "r"(*reinterpret_cast(&o32_bf16[1])), "r"(*reinterpret_cast(&o32_bf16[2])), "r"(*reinterpret_cast(&o32_bf16[3])) + : "memory"); + #pragma unroll + for (int zero_word = 0; zero_word < 4; zero_word++) { + zero32_bf16[zero_word] = 0; + } + int zero_publish_addr = (smem_inv_addr + prep_stage * 41984 + (unsigned int)((16 + lane_col) / 16 * 1024 + lane_row_1 * 32 + (16 + lane_col) % 16 * 2 ^ ((16 + lane_col) / 16 * 1024 + lane_row_1 * 32 + (16 + lane_col) % 16 * 2 >> 7 & 1) << 4)); + uint32_t _stmatrix_addr_15 = static_cast((unsigned long long)zero_publish_addr); + asm volatile("stmatrix.sync.aligned.m8n8.x4.shared.b16 [%0], {%1, %2, %3, %4};\n" + :: "r"(_stmatrix_addr_15), "r"(*reinterpret_cast(&zero32_bf16[0])), "r"(*reinterpret_cast(&zero32_bf16[1])), "r"(*reinterpret_cast(&zero32_bf16[2])), "r"(*reinterpret_cast(&zero32_bf16[3])) + : "memory"); + } else if (prep_local_warp == 1) { + int stage_f32_0_1 = prep_stage * 10496; + float restore_scale_1 = smem_restore_factor_all[stage_f32_0_1 + 128]; + float restore_factor_1[8]; + int restore_segment_1 = lane & 15; + #pragma unroll + for (int restore_elem_2 = 0; restore_elem_2 < 8; restore_elem_2++) { + int restore_col_1 = restore_segment_1 * 8 + restore_elem_2; + restore_factor_1[restore_elem_2] = smem_restore_factor_all[stage_f32_0_1 + restore_col_1]; + } + #pragma unroll 1 + for (int restore_pass_1 = 0; restore_pass_1 < 4; restore_pass_1++) { + int restore_row_1 = restore_pass_1 * 2 + (lane >> 4); + float restore_qd_values_1[8]; + float restore_kd_values_1[8]; + float restore_ki_values_1[8]; + unsigned int packed_6[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_6[0])), "=r"(*reinterpret_cast(&packed_6[(0) + 1])), "=r"(*reinterpret_cast(&packed_6[(0) + 2])), "=r"(*reinterpret_cast(&packed_6[(0) + 3])) + : "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_fp32_3[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_fp32_3[_pair * 2])[0]), "=f"((&packed_fp32_3[_pair * 2])[1]) + : "r"(packed_6[_pair + 0])); + } + #pragma unroll + for (int value_idx_6 = 0; value_idx_6 < 8; value_idx_6++) { + restore_qd_values_1[value_idx_6] = packed_fp32_3[value_idx_6]; + } + unsigned int packed_0_4[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_0_4[0])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 1])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 2])), "=r"(*reinterpret_cast(&packed_0_4[(0) + 3])) + : "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_0_fp32_2[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_0_fp32_2[_pair * 2])[0]), "=f"((&packed_0_fp32_2[_pair * 2])[1]) + : "r"(packed_0_4[_pair + 0])); + } + #pragma unroll + for (int value_idx_7 = 0; value_idx_7 < 8; value_idx_7++) { + restore_kd_values_1[value_idx_7] = packed_0_fp32_2[value_idx_7]; + } + unsigned int packed_1_3[4]; + asm volatile("ld.shared.v4.b32 {%0,%1,%2,%3}, [%4];" + : "=r"(*reinterpret_cast(&packed_1_3[0])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 1])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 2])), "=r"(*reinterpret_cast(&packed_1_3[(0) + 3])) + : "r"((smem_ki_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)))); + float packed_1_fp32_1[8]; + #pragma unroll + for (int _pair = 0; _pair < 4; _pair++) { + asm volatile( + "{\n\t" + "shl.b32 %0, %2, 16;\n\t" + "and.b32 %1, %2, 0xffff0000;\n\t" + "}\n" + : "=f"((&packed_1_fp32_1[_pair * 2])[0]), "=f"((&packed_1_fp32_1[_pair * 2])[1]) + : "r"(packed_1_3[_pair + 0])); + } + #pragma unroll + for (int value_idx_8 = 0; value_idx_8 < 8; value_idx_8++) { + restore_ki_values_1[value_idx_8] = packed_1_fp32_1[value_idx_8]; + } + float restore_kr_values_1[8]; + #pragma unroll + for (int restore_elem_3 = 0; restore_elem_3 < 8; restore_elem_3++) { + restore_kr_values_1[restore_elem_3] = restore_ki_values_1[restore_elem_3] * restore_factor_1[restore_elem_3]; + } + const float2 _scale2_16 = {restore_scale_1, restore_scale_1}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_qd_values_1)[_ls], _scale2_16); + const float2 _scale2_17 = {restore_scale_1, restore_scale_1}; + #pragma unroll + for (int _ls = 0; _ls < 4; _ls++) + mul_f32x2_inplace(&reinterpret_cast(restore_kd_values_1)[_ls], _scale2_17); + unsigned int packed_2_2[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_qd_values_1[_lp*2 + 0], restore_qd_values_1[_lp*2+1 + 0])); + packed_2_2[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_7 = 0; word_7 < 4; word_7++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_qd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_7 * 4)), "r"(packed_2_2[word_7])); + } + unsigned int packed_3_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kd_values_1[_lp*2 + 0], restore_kd_values_1[_lp*2+1 + 0])); + packed_3_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_8 = 0; word_8 < 4; word_8++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kd_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_8 * 4)), "r"(packed_3_1[word_8])); + } + unsigned int packed_4_1[4]; + #pragma unroll + for (int _lp = 0; _lp < 4; _lp++) { + __nv_bfloat162 _bf2 = __float22bfloat162_rn(make_float2(restore_kr_values_1[_lp*2 + 0], restore_kr_values_1[_lp*2+1 + 0])); + packed_4_1[_lp] = *(uint32_t*)&_bf2; + } + #pragma unroll + for (int word_9 = 0; word_9 < 4; word_9++) { + asm volatile("st.shared.b32 [%0], %1;" :: "r"((smem_kr_trans_addr + prep_stage * 41984 + (unsigned int)(restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 ^ (restore_segment_1 * 8 / 64 * 4096 + restore_row_1 * 128 + restore_segment_1 * 8 % 64 * 2 >> 7 & 7) << 4)) + (unsigned int)(word_9 * 4)), "r"(packed_4_1[word_9])); + } + } + } + asm volatile("fence.proxy.async.shared::cta;" ::: "memory"); + if (prep_instance == 0) { + asm volatile("barrier.sync 11, 128;" ::: "memory"); + } else if (prep_instance == 1) { + asm volatile("barrier.sync 12, 128;" ::: "memory"); + } else { + if (prep_instance == 2) { + asm volatile("barrier.sync 13, 128;" ::: "memory"); + } else if (prep_instance == 3) { + asm volatile("barrier.sync 14, 128;" ::: "memory"); + } else { + asm volatile("barrier.sync 15, 128;" ::: "memory"); + } + } + long long packet_idx = + (long long)task_idx_4 * mailbox_depth + + chunk_idx_3 % mailbox_depth; + if (prep_tid == 0) { + const unsigned int generation = + (unsigned int)(chunk_idx_3 / mailbox_depth); + const unsigned int reusable_state = generation == 0 + ? 0u + : ((generation - 1u) << 3) | 7u; + wait_k1_global_flag( + k1_flags + packet_idx, reusable_state); + } + const unsigned int ready_value = + ((unsigned int)(chunk_idx_3 / mailbox_depth) << 3) | 1u; + publish_k1_to_global( + smem_qd_addr + prep_stage * 41984, + k1_workspace + packet_idx * kK1PacketBytes, + k1_flags + packet_idx, ready_value, prep_tid); + if (prep_tid == 0) { + mbarrier_arrive( + raw_inputs_free_addr + prep_stage * 8); + mbarrier_arrive(smem_free_addr + prep_stage * 8); + } + for (int _advance = 0; _advance < 5; _advance++) { + prep_stage += 1; + if (prep_stage == 5) { prep_stage = 0; _phase_raw_inputs_free ^= 1; _phase_smem_free ^= 1; _phase_gate_raw_full ^= 1; _phase_qk_raw_full ^= 1; _phase_prep_diag_ready ^= 1; _phase_prep_inv16_ready ^= 1; } + } + } + } + } + + __syncthreads(); + cluster_sync(); + +} + +} // extern "C" + +// clang-format on diff --git a/csrc/kda/flashkda_bf16_fused_m64_k1_parallel_binding.cu b/csrc/kda/flashkda_bf16_fused_m64_k1_parallel_binding.cu new file mode 100644 index 00000000000..24a28e46432 --- /dev/null +++ b/csrc/kda/flashkda_bf16_fused_m64_k1_parallel_binding.cu @@ -0,0 +1,159 @@ +/* + * Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + */ + +#include "flashkda_binding_common.cuh" + +#define uint8_t flashkda_m64_k1_parallel_uint8_t +#define uint16_t flashkda_m64_k1_parallel_uint16_t +#define uint32_t flashkda_m64_k1_parallel_uint32_t +#define uint64_t flashkda_m64_k1_parallel_uint64_t +#define int32_t flashkda_m64_k1_parallel_int32_t +#define int16_t flashkda_m64_k1_parallel_int16_t +#include "flashkda_bf16_fused_m64_k1_parallel.cu" +#undef uint8_t +#undef uint16_t +#undef uint32_t +#undef uint64_t +#undef int32_t +#undef int16_t + +namespace flashinfer { +namespace flash_kda { + +constexpr int64_t kK1ParallelPacketBytes = 31520; +static_assert(kK1ParallelPacketBytes == kK1PacketBytes); +static_assert(THREADS == 1024); +static_assert(SMEM_TOTAL == 219136); + +void RunM64K1Parallel(TensorView q, TensorView k, TensorView v, TensorView g, TensorView beta, + TensorView beta_tma, TensorView A_log, TensorView dt_bias, + TensorView cu_seqlens, TensorView seq_order, TensorView initial_state, + TensorView out, TensorView final_state, TensorView descriptor_storage, + TensorView k1_workspace, int64_t prepare_descriptors, int64_t num_heads, + int64_t use_initial_state, int64_t store_final_state, int64_t cluster_size, + int64_t mailbox_depth, double scale, double lower_bound, + int64_t cuda_stream) { + TVM_FFI_ICHECK(cuda_stream >= 0) << "cuda_stream must be a non-negative stream handle"; + TVM_FFI_ICHECK(q.device().device_type == kDLCUDA) << "q must be a CUDA tensor"; + const int32_t device_id = q.device().device_id; + ffi::CUDADeviceGuard device_guard(device_id); + CheckFlashKDATarget(device_id); + + const int64_t num_seqs = + CheckCommonInputs(q, k, v, g, beta, beta_tma, A_log, dt_bias, cu_seqlens, seq_order, + initial_state, out, final_state, descriptor_storage, prepare_descriptors, + num_heads, use_initial_state, store_final_state, scale, lower_bound); + CheckCudaTensor(k1_workspace, "k1_workspace", device_id); + CheckDtype(k1_workspace, "k1_workspace", dl_uint8); + for (const auto& named : { + std::pair(&q, "q"), + std::pair(&k, "k"), + std::pair(&v, "v"), + std::pair(&g, "g"), + std::pair(&beta, "beta"), + std::pair(&beta_tma, "beta_tma"), + std::pair(&A_log, "A_log"), + std::pair(&dt_bias, "dt_bias"), + std::pair(&cu_seqlens, "cu_seqlens"), + std::pair(&seq_order, "seq_order"), + std::pair(&out, "out"), + std::pair(&descriptor_storage, "descriptor_storage"), + }) { + CheckNoOverlap(k1_workspace, "k1_workspace", *named.first, named.second); + } + if (use_initial_state != 0) { + CheckNoOverlap(k1_workspace, "k1_workspace", initial_state, "initial_state"); + } + if (store_final_state != 0) { + CheckNoOverlap(k1_workspace, "k1_workspace", final_state, "final_state"); + } + TVM_FFI_ICHECK(cluster_size == 4) << "M64 K1-parallel FlashKDA requires cluster_size == 4"; + TVM_FFI_ICHECK(mailbox_depth > 0 && mailbox_depth <= std::numeric_limits::max()) + << "mailbox_depth must be in the positive int32 range"; + const int64_t producer_instances = (cluster_size - 2) * 5; + TVM_FFI_ICHECK(mailbox_depth >= producer_instances && mailbox_depth % producer_instances == 0) + << "mailbox_depth must be a positive multiple of the helper producer count " + << producer_instances << " for C" << cluster_size << "; got " << mailbox_depth; + TVM_FFI_ICHECK(SupportsBetaTmaHeadCount(num_heads)) + << "K1-parallel FlashKDA requires H == 1, H == 4, or H >= 8 and divisible by 8"; + + const int64_t num_tasks = num_seqs * num_heads; + const int64_t packet_count = num_tasks * mailbox_depth; + TVM_FFI_ICHECK(packet_count > 0 && + packet_count <= std::numeric_limits::max() / kK1ParallelPacketBytes) + << "K1 mailbox packet count is out of range"; + const int64_t flag_offset = + (packet_count * kK1ParallelPacketBytes + int64_t{255}) & ~int64_t{255}; + TVM_FFI_ICHECK(packet_count <= (std::numeric_limits::max() - flag_offset) / + static_cast(sizeof(uint32_t))) + << "K1 mailbox flag size is out of range"; + const int64_t required_bytes = + flag_offset + packet_count * static_cast(sizeof(uint32_t)); + TVM_FFI_ICHECK(k1_workspace.numel() >= required_bytes) + << "k1_workspace requires " << required_bytes << " bytes, got " << k1_workspace.numel(); + TVM_FFI_ICHECK(reinterpret_cast(k1_workspace.data_ptr()) % 256 == 0) + << "k1_workspace must be 256-byte aligned"; + + constexpr int32_t kSmemBytes = SMEM_TOTAL; + CheckDynamicSmemCapacity(device_id, kSmemBytes); + CheckCuda(cudaFuncSetAttribute(kernel_flashkda_bf16_fused_m64_k1_parallel, + cudaFuncAttributeMaxDynamicSharedMemorySize, kSmemBytes), + "cudaFuncSetAttribute(kernel_flashkda_bf16_fused_m64_k1_parallel)"); + + const cudaStream_t stream = reinterpret_cast(static_cast(cuda_stream)); + const TmaPointers tma = EncodeTmaPointers<64>(q, k, v, g, beta_tma, out, descriptor_storage, + prepare_descriptors, stream); + PackBetaForTmaIfNeeded(beta, beta_tma, num_heads, stream); + + auto* workspace_bytes = static_cast(k1_workspace.data_ptr()); + auto* flags = reinterpret_cast(workspace_bytes + flag_offset); + CheckCuda(cudaMemsetAsync(flags, 0, packet_count * sizeof(uint32_t), stream), + "cudaMemsetAsync(K1 mailbox flags)"); + + const int64_t grid_x_i64 = num_tasks * cluster_size; + TVM_FFI_ICHECK(grid_x_i64 > 0 && grid_x_i64 <= std::numeric_limits::max()) + << "K1-parallel FlashKDA grid.x is out of range: " << grid_x_i64; + + cudaLaunchAttribute attribute{}; + attribute.id = cudaLaunchAttributeClusterDimension; + attribute.val.clusterDim = {static_cast(cluster_size), 1u, 1u}; + cudaLaunchConfig_t config{}; + config.gridDim = dim3(static_cast(grid_x_i64), 1, 1); + config.blockDim = dim3(THREADS, 1, 1); + config.dynamicSmemBytes = kSmemBytes; + config.stream = stream; + config.attrs = &attribute; + config.numAttrs = 1; + + CheckCuda(cudaLaunchKernelEx(&config, kernel_flashkda_bf16_fused_m64_k1_parallel, + reinterpret_cast<__nv_bfloat16*>(q.data_ptr()), tma.q, + reinterpret_cast<__nv_bfloat16*>(k.data_ptr()), tma.k, + reinterpret_cast<__nv_bfloat16*>(v.data_ptr()), tma.v, + reinterpret_cast<__nv_bfloat16*>(g.data_ptr()), tma.g, + reinterpret_cast<__nv_bfloat16*>(beta.data_ptr()), tma.beta, + reinterpret_cast(A_log.data_ptr()), + reinterpret_cast(dt_bias.data_ptr()), + reinterpret_cast(cu_seqlens.data_ptr()), + reinterpret_cast(seq_order.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(initial_state.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), tma.out, + reinterpret_cast<__nv_bfloat16*>(final_state.data_ptr()), + workspace_bytes, flags, static_cast(mailbox_depth), + static_cast(cluster_size), static_cast(num_heads), + static_cast(use_initial_state), + static_cast(store_final_state), static_cast(scale), + static_cast(lower_bound)), + "kernel_flashkda_bf16_fused_m64_k1_parallel launch"); +} + +} // namespace flash_kda +} // namespace flashinfer + +TVM_FFI_DLL_EXPORT_TYPED_FUNC(run, flashinfer::flash_kda::RunM64K1Parallel); diff --git a/csrc/kda/flashkda_binding_common.cuh b/csrc/kda/flashkda_binding_common.cuh index bb9b82b5ccb..debd0e7b506 100644 --- a/csrc/kda/flashkda_binding_common.cuh +++ b/csrc/kda/flashkda_binding_common.cuh @@ -48,6 +48,11 @@ inline int64_t RoundUpBetaTmaHeads(int64_t num_heads) { kBetaTmaHeadsPerBox; } +inline bool SupportsBetaTmaHeadCount(int64_t num_heads) { + return num_heads == 1 || num_heads == 4 || + (num_heads >= 8 && num_heads % kBetaTmaHeadsPerBox == 0); +} + static __global__ void PackBetaForTmaKernel(const __nv_bfloat16* beta, __nv_bfloat16* beta_tma, int64_t token_count, int64_t padded_elements, int64_t num_heads, int64_t padded_num_heads) { diff --git a/docs/_static/cake-kda-small-bh-k1-parallelism.webp b/docs/_static/cake-kda-small-bh-k1-parallelism.webp new file mode 100644 index 00000000000..ad7a2aa23b4 Binary files /dev/null and b/docs/_static/cake-kda-small-bh-k1-parallelism.webp differ diff --git a/flashinfer/aot.py b/flashinfer/aot.py index 2fcc6fc1684..a70c1d82fcf 100644 --- a/flashinfer/aot.py +++ b/flashinfer/aot.py @@ -64,7 +64,9 @@ from .jit.flash_kda import ( FlashKDATarget, gen_flash_kda_m64_module, + gen_flash_kda_m64_k1_parallel_module, gen_flash_kda_m128_module, + gen_flash_kda_m128_k1_parallel_module, ) from .jit.flash_kda_decode import ( FLASH_KDA_DECODE_DIRECT_VARIANTS, @@ -558,7 +560,9 @@ def gen_all_modules( jit_specs.extend( [ gen_flash_kda_m64_module(flash_kda_target), + gen_flash_kda_m64_k1_parallel_module(flash_kda_target), gen_flash_kda_m128_module(flash_kda_target), + gen_flash_kda_m128_k1_parallel_module(flash_kda_target), ] ) diff --git a/flashinfer/jit/__init__.py b/flashinfer/jit/__init__.py index 81a07730382..59c409b8112 100644 --- a/flashinfer/jit/__init__.py +++ b/flashinfer/jit/__init__.py @@ -105,18 +105,30 @@ from .flash_kda import ( gen_flash_kda_m64_module as gen_flash_kda_m64_module, ) +from .flash_kda import ( + gen_flash_kda_m64_k1_parallel_module as gen_flash_kda_m64_k1_parallel_module, +) from .flash_kda import ( gen_flash_kda_m128_module as gen_flash_kda_m128_module, ) +from .flash_kda import ( + gen_flash_kda_m128_k1_parallel_module as gen_flash_kda_m128_k1_parallel_module, +) from .flash_kda import ( get_flash_kda_prefill_module as get_flash_kda_prefill_module, ) from .flash_kda import ( load_flash_kda_m64_module as load_flash_kda_m64_module, ) +from .flash_kda import ( + load_flash_kda_m64_k1_parallel_module as load_flash_kda_m64_k1_parallel_module, +) from .flash_kda import ( load_flash_kda_m128_module as load_flash_kda_m128_module, ) +from .flash_kda import ( + load_flash_kda_m128_k1_parallel_module as load_flash_kda_m128_k1_parallel_module, +) from .nvfp4_attention_sm120 import ( gen_nvfp4_attention_sm120_module as gen_nvfp4_attention_sm120_module, ) diff --git a/flashinfer/jit/flash_kda.py b/flashinfer/jit/flash_kda.py index 15ea0608e3e..c373668073e 100644 --- a/flashinfer/jit/flash_kda.py +++ b/flashinfer/jit/flash_kda.py @@ -27,7 +27,7 @@ sm100f_nvcc_flags, ) -FlashKDAVariant = Literal["m64", "m128"] +FlashKDAVariant = Literal["m64", "m128", "m64_k1_parallel", "m128_k1_parallel"] FlashKDATarget = Literal["sm100a", "sm100f"] _FLASH_KDA_NVCC_FLAGS = { @@ -76,7 +76,7 @@ def _get_flash_kda_include_dir() -> Path: def get_flash_kda_uri(variant: FlashKDAVariant, target: FlashKDATarget) -> str: """Return the target-specific JIT/AOT key for one schedule.""" - if variant not in ("m64", "m128"): + if variant not in ("m64", "m128", "m64_k1_parallel", "m128_k1_parallel"): raise ValueError(f"unsupported FlashKDA variant: {variant}") if target not in _FLASH_KDA_NVCC_FLAGS: raise ValueError(f"unsupported FlashKDA target: {target}") @@ -119,7 +119,7 @@ def gen_flash_kda_module(variant: FlashKDAVariant, target: FlashKDATarget) -> Ji def gen_flash_kda_m64_module(target: FlashKDATarget) -> JitSpec: - """Generate the fixed N=1, H=64 two-CTA M64 module.""" + """Generate the two-CTA-per-task M64 module.""" return gen_flash_kda_module("m64", target) @@ -130,6 +130,18 @@ def gen_flash_kda_m128_module(target: FlashKDATarget) -> JitSpec: return gen_flash_kda_module("m128", target) +def gen_flash_kda_m128_k1_parallel_module(target: FlashKDATarget) -> JitSpec: + """Generate the SM100-family owner/helper M128 module.""" + + return gen_flash_kda_module("m128_k1_parallel", target) + + +def gen_flash_kda_m64_k1_parallel_module(target: FlashKDATarget) -> JitSpec: + """Generate the SM100-family dual-owner M64 module.""" + + return gen_flash_kda_module("m64_k1_parallel", target) + + @functools.cache def load_flash_kda_module(variant: FlashKDAVariant, target: FlashKDATarget): """Build or load one physical, target-specific FlashKDA module.""" @@ -140,7 +152,7 @@ def load_flash_kda_module(variant: FlashKDAVariant, target: FlashKDATarget): def load_flash_kda_m64_module(target: FlashKDATarget): - """Load the fixed N=1, H=64 two-CTA M64 module.""" + """Load the two-CTA-per-task M64 module.""" return load_flash_kda_module("m64", target) @@ -151,6 +163,18 @@ def load_flash_kda_m128_module(target: FlashKDATarget): return load_flash_kda_module("m128", target) +def load_flash_kda_m128_k1_parallel_module(target: FlashKDATarget): + """Load the SM100-family owner/helper M128 module.""" + + return load_flash_kda_module("m128_k1_parallel", target) + + +def load_flash_kda_m64_k1_parallel_module(target: FlashKDATarget): + """Load the SM100-family dual-owner M64 module.""" + + return load_flash_kda_module("m64_k1_parallel", target) + + def get_flash_kda_prefill_module(variant: FlashKDAVariant, target: FlashKDATarget): """Return the loaded module used by the recurrent-KDA prefill dispatcher.""" @@ -161,11 +185,15 @@ def get_flash_kda_prefill_module(variant: FlashKDAVariant, target: FlashKDATarge "FlashKDATarget", "FlashKDAVariant", "gen_flash_kda_m64_module", + "gen_flash_kda_m64_k1_parallel_module", "gen_flash_kda_m128_module", + "gen_flash_kda_m128_k1_parallel_module", "gen_flash_kda_module", "get_flash_kda_prefill_module", "get_flash_kda_uri", "load_flash_kda_m64_module", + "load_flash_kda_m64_k1_parallel_module", "load_flash_kda_m128_module", + "load_flash_kda_m128_k1_parallel_module", "load_flash_kda_module", ] diff --git a/flashinfer/kda_prefill.py b/flashinfer/kda_prefill.py index 1e5b69f6d99..7300d1182f2 100644 --- a/flashinfer/kda_prefill.py +++ b/flashinfer/kda_prefill.py @@ -38,10 +38,19 @@ _FLASH_KDA_BETA_TMA_HEADS_PER_BOX = 8 _FLASH_KDA_SUPPORTED_COMPUTE_CAPABILITIES = {(10, 0), (10, 3)} _FLASH_KDA_DESCRIPTOR_STORAGE_BYTES = 6 * 128 +_FLASH_KDA_K1_PACKET_BYTES = 31_520 +_FLASH_KDA_K1_PARALLEL_VARIANTS = ( + "m64_k1_parallel", + "m128_k1_parallel", +) _flash_kda_tensor_cache: dict[tuple, torch.Tensor] = {} _flash_kda_tensor_cache_lock = threading.Lock() +def _flash_kda_head_count_supports_tma(num_heads: int) -> bool: + return num_heads in (1, 4) or (num_heads >= 8 and num_heads % 8 == 0) + + class _RecurrentKDAPrefillWorkspaceBase: def __init__(self, device: torch.device | str) -> None: normalized_device = torch.device(device) @@ -53,13 +62,14 @@ def __init__(self, device: torch.device | str) -> None: self._lock = threading.Lock() self._state_scratch: Optional[torch.Tensor] = None self._beta_padding: Optional[torch.Tensor] = None + self._k1_mailbox: Optional[torch.Tensor] = None self._descriptor_storages = { variant: torch.empty( _FLASH_KDA_DESCRIPTOR_STORAGE_BYTES, dtype=torch.uint8, device=self.device, ) - for variant in ("m64", "m128") + for variant in ("m64", "m128", *_FLASH_KDA_K1_PARALLEL_VARIANTS) } self._descriptor_signatures: dict[str, tuple] = {} self._bound_stream_ptr: Optional[int] = None @@ -74,8 +84,8 @@ class RecurrentKDAPrefillWorkspace(_RecurrentKDAPrefillWorkspaceBase): device. Warm it by invoking that function eagerly with the exact tensors and capture stream, then synchronize that stream before capture. The workspace owns optional final-state scratch for calls without an initial - state, beta padding, and M64/M128 TMA descriptor storage for the lifetime - of the graph. + state, beta padding, K1 owner/helper mailbox storage, and M64/M128 TMA + descriptor storage for the lifetime of the graph. A workspace binds to its first stream. Once it participates in capture it cannot be passed to Python again, either eagerly or in another capture. @@ -243,11 +253,86 @@ def _flash_kda_prefill_is_eligible( def _select_flash_kda_prefill_variant( - *, fixed_layout: bool, num_sequences: int, num_heads: int -) -> "FlashKDAVariant": - if fixed_layout and num_sequences == 1 and num_heads == 64: - return "m64" - return "m128" + *, + fixed_layout: bool, + num_sequences: int, + num_heads: int, + sequence_length: int, + device: torch.device, +) -> tuple["FlashKDAVariant", int, int]: + """Select the measured SM100-family oracle and K1-helper schedule. + + ``sequence_length`` is the per-sequence length for fixed input and the + packed token count for varlen input. Packed dispatch deliberately uses the + host-known average length so it never reads ``cu_seqlens`` back to the CPU. + """ + + compute_capability = get_compute_capability(device) + if compute_capability not in _FLASH_KDA_SUPPORTED_COMPUTE_CAPABILITIES: + if fixed_layout and num_sequences == 1 and num_heads == 64: + return "m64", 0, 0 + return "m128", 0, 0 + + task_count = num_sequences * num_heads + average_sequence_length = ( + sequence_length if fixed_layout else sequence_length // num_sequences + ) + helpers_supported = fixed_layout or compute_capability == (10, 0) + if num_heads == 1 and compute_capability != (10, 0): + return "m128", 0, 0 + if ( + compute_capability == (10, 0) + and fixed_layout + and num_sequences == 1 + and num_heads == 1 + and average_sequence_length >= 4096 + ): + return "m64_k1_parallel", 4, 10 + if ( + helpers_supported + and _flash_kda_head_count_supports_tma(num_heads) + and (num_heads != 1 or fixed_layout) + and average_sequence_length >= 2048 + ): + # Keep only schedules that won the architecture-specific cold-L2 + # forced-route sweeps used to tune this dispatch. + depth_schedule = ( + ((8, 15), (32, 30)) + if compute_capability == (10, 0) + else ((8, 30), (32, 45)) + ) + for max_tasks, mailbox_depth in depth_schedule: + if task_count <= max_tasks: + return "m128_k1_parallel", 4, mailbox_depth + if not fixed_layout: + return "m128", 0, 0 + if num_sequences == 1 and num_heads == 64: + return "m64", 0, 0 + return "m128", 0, 0 + + +def _k1_mailbox_bytes(task_count: int, mailbox_depth: int) -> int: + packet_count = task_count * mailbox_depth + flag_offset = (packet_count * _FLASH_KDA_K1_PACKET_BYTES + 255) & ~255 + return flag_offset + packet_count * 4 + + +def _k1_mailbox_workspace( + *, + workspace: _RecurrentKDAPrefillWorkspaceBase, + device: torch.device, + required_bytes: int, +) -> torch.Tensor: + mailbox = workspace._k1_mailbox + if mailbox is None or mailbox.numel() < required_bytes: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "recurrent_kda K1 mailbox is not large enough for CUDA graph " + "capture; warm the largest owner/helper shape first" + ) + mailbox = torch.empty(required_bytes, dtype=torch.uint8, device=device) + workspace._k1_mailbox = mailbox + return mailbox[:required_bytes] def _cached_tensor( @@ -665,10 +750,12 @@ def _run_flash_kda_prefill( ) if not math.isfinite(scale_value): raise ValueError(f"scale must be finite, got {scale_value}") - variant = _select_flash_kda_prefill_variant( + variant, cluster_size, mailbox_depth = _select_flash_kda_prefill_variant( fixed_layout=fixed_layout, num_sequences=num_sequences, num_heads=num_heads, + sequence_length=seq_len, + device=q.device, ) target = _select_flash_kda_prefill_target(q.device) stream_ptr = int(torch.cuda.current_stream(q.device).cuda_stream) @@ -720,7 +807,7 @@ def _run_flash_kda_prefill( descriptor_storage = workspace._descriptor_storages[variant] module = _get_flash_kda_prefill_module(variant, target) try: - module.run( + common_args = ( q, k, v, @@ -735,14 +822,39 @@ def _run_flash_kda_prefill( out_buf, final_state_arg, descriptor_storage, - prepare_descriptors, - num_heads, - int(use_initial_state), - int(store_final_state), - scale_value, - float(lower_bound), - stream_ptr, ) + if variant in _FLASH_KDA_K1_PARALLEL_VARIANTS: + k1_mailbox = _k1_mailbox_workspace( + workspace=workspace, + device=q.device, + required_bytes=_k1_mailbox_bytes( + num_sequences * num_heads, mailbox_depth + ), + ) + module.run( + *common_args, + k1_mailbox, + prepare_descriptors, + num_heads, + int(use_initial_state), + int(store_final_state), + cluster_size, + mailbox_depth, + scale_value, + float(lower_bound), + stream_ptr, + ) + else: + module.run( + *common_args, + prepare_descriptors, + num_heads, + int(use_initial_state), + int(store_final_state), + scale_value, + float(lower_bound), + stream_ptr, + ) except Exception: if prepare_descriptors: workspace._descriptor_signatures.pop(variant, None) diff --git a/tests/jit/test_flash_kda_jit.py b/tests/jit/test_flash_kda_jit.py index 0b0f07b7be5..85c0a246a0c 100644 --- a/tests/jit/test_flash_kda_jit.py +++ b/tests/jit/test_flash_kda_jit.py @@ -207,6 +207,71 @@ def test_flash_kda_descriptor_workspace_contract(): ) +@pytest.mark.parametrize("target", ["sm100a", "sm100f"]) +def test_flash_kda_k1_parallel_jit_spec(monkeypatch, target): + target_arch = (10, "0a") if target == "sm100a" else (10, "0f") + monkeypatch.setattr( + jit_core.current_compilation_context, + "TARGET_CUDA_ARCHS", + {target_arch}, + ) + flash_kda.gen_flash_kda_module.cache_clear() + + spec = flash_kda.gen_flash_kda_m128_k1_parallel_module(target) + + assert spec.name == f"flash_kda_bf16_fused_m128_k1_parallel_{target}" + assert [source.name for source in spec.sources] == [ + "flashkda_bf16_fused_m128_k1_parallel_binding.cu" + ] + source = spec.sources[0].parent / "flashkda_bf16_fused_m128_k1_parallel.cu" + text = source.read_text() + assert hashlib.sha256(text.encode()).hexdigest() == ( + "b530ff5593658f28de68845f5c633d156f375ee8452d0e65634c7588bf369782" + ) + assert "bounded global-mailbox packets" in text + assert "constexpr int kK1PacketBytes = 31520;" in text + assert "wait_k1_global_flag" in text + + binding_text = spec.sources[0].read_text() + assert "CheckFlashKDATarget(device_id);" in binding_text + assert "cluster_size == 4 || cluster_size == 8" in binding_text + assert "mailbox_depth % producer_instances == 0" in binding_text + assert "cluster_size < 0" not in binding_text + assert "global_pool" not in text + + +@pytest.mark.parametrize("target", ["sm100a", "sm100f"]) +def test_flash_kda_m64_k1_parallel_jit_spec(monkeypatch, target): + target_arch = (10, "0a") if target == "sm100a" else (10, "0f") + monkeypatch.setattr( + jit_core.current_compilation_context, + "TARGET_CUDA_ARCHS", + {target_arch}, + ) + flash_kda.gen_flash_kda_module.cache_clear() + + spec = flash_kda.gen_flash_kda_m64_k1_parallel_module(target) + + assert spec.name == f"flash_kda_bf16_fused_m64_k1_parallel_{target}" + assert [source.name for source in spec.sources] == [ + "flashkda_bf16_fused_m64_k1_parallel_binding.cu" + ] + source = spec.sources[0].parent / "flashkda_bf16_fused_m64_k1_parallel.cu" + text = source.read_text() + assert hashlib.sha256(text.encode()).hexdigest() == ( + "0fd0e877cf3084423d4deb1b135164380ac2bcbea99366792918db060c1f0c91" + ) + assert "constexpr int kOwnerCount = 2;" in text + assert "(generation << 3) | 1u" in text + assert "acknowledge_k1_global_packet" in text + + binding_text = spec.sources[0].read_text() + assert "(cluster_size - 2) * 5" in binding_text + assert "EncodeTmaPointers<64>" in binding_text + assert "cluster_size == 4" in binding_text + assert "cluster_size" in text + + def test_flash_kda_variant_validation_and_public_getter(monkeypatch): with pytest.raises(ValueError, match="unsupported FlashKDA variant"): flash_kda.get_flash_kda_uri("m32", "sm100f") @@ -292,7 +357,7 @@ def get_nvcc_flags_list(self, supported_major_versions=None): ({"flash_kda_prefill_sm100f": True}, "sm100f"), ], ) -def test_aot_registers_two_flash_kda_modules( +def test_aot_registers_all_flash_kda_modules( monkeypatch, capabilities, expected_target ): from flashinfer import aot @@ -313,6 +378,16 @@ def fake_flash_kda(variant, target): "gen_flash_kda_m128_module", lambda target: fake_flash_kda("m128", target), ) + monkeypatch.setattr( + aot, + "gen_flash_kda_m128_k1_parallel_module", + lambda target: fake_flash_kda("m128_k1_parallel", target), + ) + monkeypatch.setattr( + aot, + "gen_flash_kda_m64_k1_parallel_module", + lambda target: fake_flash_kda("m64_k1_parallel", target), + ) monkeypatch.setattr( aot, "gen_spdlog_module", lambda: SimpleNamespace(name="spdlog") ) @@ -340,11 +415,15 @@ def fake_flash_kda(variant, target): assert calls == [ ("m64", expected_target), + ("m64_k1_parallel", expected_target), ("m128", expected_target), + ("m128_k1_parallel", expected_target), ] assert [spec.name for spec in specs] == [ "spdlog", f"flash_kda_m64_{expected_target}", + f"flash_kda_m64_k1_parallel_{expected_target}", f"flash_kda_m128_{expected_target}", + f"flash_kda_m128_k1_parallel_{expected_target}", "cudnn", ] diff --git a/tests/kda/test_recurrent_kda_prefill.py b/tests/kda/test_recurrent_kda_prefill.py index 76ff26b1268..9154823fcb9 100644 --- a/tests/kda/test_recurrent_kda_prefill.py +++ b/tests/kda/test_recurrent_kda_prefill.py @@ -281,6 +281,7 @@ def test_multi_token_gqa_stays_on_existing_backend(cuda_device, monkeypatch): @pytest.mark.parametrize( ("packed", "num_heads", "expected_variant"), [ + (False, 4, "m128"), (False, 64, "m64"), (True, 64, "m128"), (True, 2, "m128"), @@ -362,6 +363,195 @@ def get_module(variant, target): assert args[5].data_ptr() != inputs["beta"].data_ptr() +@pytest.mark.parametrize( + ("num_sequences", "num_heads", "sequence_length", "expected"), + [ + (1, 1, 512, ("m128", 0, 0)), + (1, 1, 1024, ("m128", 0, 0)), + (1, 1, 2048, ("m128_k1_parallel", 4, 15)), + (1, 1, 4096, ("m64_k1_parallel", 4, 10)), + (2, 1, 2048, ("m128_k1_parallel", 4, 15)), + (1, 4, 1024, ("m128", 0, 0)), + (1, 4, 2048, ("m128_k1_parallel", 4, 15)), + (1, 8, 1024, ("m128", 0, 0)), + (1, 8, 2048, ("m128_k1_parallel", 4, 15)), + (1, 32, 4096, ("m128_k1_parallel", 4, 30)), + (1, 48, 4096, ("m128", 0, 0)), + (1, 64, 8192, ("m64", 0, 0)), + (1, 64, 4096, ("m64", 0, 0)), + (1, 72, 8192, ("m128", 0, 0)), + (1, 80, 8192, ("m128", 0, 0)), + (2, 16, 2048, ("m128_k1_parallel", 4, 30)), + ], +) +def test_k1_parallel_b200_oracle( + monkeypatch, num_sequences, num_heads, sequence_length, expected +): + monkeypatch.setattr( + kda_prefill_api, "get_compute_capability", lambda device: (10, 0) + ) + actual = kda_prefill_api._select_flash_kda_prefill_variant( + fixed_layout=True, + num_sequences=num_sequences, + num_heads=num_heads, + sequence_length=sequence_length, + device=torch.device("cuda:0"), + ) + + assert actual == expected + + +@pytest.mark.parametrize( + ("num_sequences", "num_heads", "sequence_length", "expected"), + [ + (1, 1, 4096, ("m128", 0, 0)), + (1, 4, 1024, ("m128", 0, 0)), + (1, 4, 2048, ("m128_k1_parallel", 4, 30)), + (1, 8, 1024, ("m128", 0, 0)), + (1, 8, 2048, ("m128_k1_parallel", 4, 30)), + (1, 16, 4096, ("m128_k1_parallel", 4, 45)), + (1, 32, 8192, ("m128_k1_parallel", 4, 45)), + (2, 8, 4096, ("m128_k1_parallel", 4, 45)), + (4, 8, 8192, ("m128_k1_parallel", 4, 45)), + (1, 48, 4096, ("m128", 0, 0)), + (1, 80, 8192, ("m128", 0, 0)), + ], +) +def test_k1_parallel_b300_oracle( + monkeypatch, num_sequences, num_heads, sequence_length, expected +): + monkeypatch.setattr( + kda_prefill_api, "get_compute_capability", lambda device: (10, 3) + ) + actual = kda_prefill_api._select_flash_kda_prefill_variant( + fixed_layout=True, + num_sequences=num_sequences, + num_heads=num_heads, + sequence_length=sequence_length, + device=torch.device("cuda:0"), + ) + assert actual == expected + + +@pytest.mark.parametrize( + ("compute_capability", "num_sequences", "num_heads", "total_tokens", "expected"), + [ + ((10, 0), 1, 4, 8192, ("m128_k1_parallel", 4, 15)), + ((10, 3), 1, 4, 8192, ("m128", 0, 0)), + ((10, 0), 1, 8, 8192, ("m128_k1_parallel", 4, 15)), + ((10, 0), 2, 8, 4096, ("m128_k1_parallel", 4, 30)), + ((10, 3), 1, 8, 8192, ("m128", 0, 0)), + ((10, 3), 2, 8, 4096, ("m128", 0, 0)), + ((10, 3), 4, 8, 4096, ("m128", 0, 0)), + ((10, 3), 4, 8, 8192, ("m128", 0, 0)), + ((10, 3), 5, 8, 10240, ("m128", 0, 0)), + ((10, 3), 1, 12, 8192, ("m128", 0, 0)), + ], +) +def test_k1_parallel_varlen_oracle( + monkeypatch, + compute_capability, + num_sequences, + num_heads, + total_tokens, + expected, +): + monkeypatch.setattr( + kda_prefill_api, + "get_compute_capability", + lambda device: compute_capability, + ) + assert ( + kda_prefill_api._select_flash_kda_prefill_variant( + fixed_layout=False, + num_sequences=num_sequences, + num_heads=num_heads, + sequence_length=total_tokens, + device=torch.device("cuda:0"), + ) + == expected + ) + + +def test_k1_parallel_varlen_fallback_stays_m128(monkeypatch): + monkeypatch.setattr( + kda_prefill_api, "get_compute_capability", lambda device: (10, 0) + ) + assert kda_prefill_api._select_flash_kda_prefill_variant( + fixed_layout=False, + num_sequences=1, + num_heads=64, + sequence_length=1024, + device=torch.device("cuda:0"), + ) == ("m128", 0, 0) + + +def test_k1_mailbox_size_is_bounded(): + assert kda_prefill_api._k1_mailbox_bytes(8, 15) == 3_782_880 + assert kda_prefill_api._k1_mailbox_bytes(32, 30) == 30_263_040 + assert kda_prefill_api._k1_mailbox_bytes(32, 45) == 45_394_560 + assert kda_prefill_api._k1_mailbox_bytes(32, 45) < 48_000_000 + + +def test_k1_parallel_route_and_ffi_abi(cuda_device, monkeypatch): + monkeypatch.setattr( + kda_prefill_api, "get_compute_capability", lambda device: (10, 0) + ) + monkeypatch.setattr( + kda_prefill_api, "_is_cuda_version_at_least", lambda version: True + ) + monkeypatch.setattr(kda_prefill_api, "_flash_kda_stream_workspaces", {}) + module = _RecorderModule() + routes = [] + + def get_module(variant, target): + routes.append((variant, target)) + return module + + monkeypatch.setattr(kda_prefill_api, "_get_flash_kda_prefill_module", get_module) + inputs = _make_inputs(seq_lens=[2048], num_heads=8, packed=False) + recurrent_kda(**_strict_prefill_kwargs(inputs)) + + assert routes == [("m128_k1_parallel", "sm100f")] + (args,) = module.calls + assert len(args) == 24 + assert args[13].shape == (768,) + assert args[14].dtype == torch.uint8 + assert args[14].numel() == kda_prefill_api._k1_mailbox_bytes(8, 15) + assert args[15] == 1 + assert args[16:21] == (8, 0, 0, 4, 15) + assert math.isclose(args[21], 128**-0.5) + assert args[22] == -5.0 + assert args[23] == int(torch.cuda.current_stream(cuda_device).cuda_stream) + + +@pytest.mark.parametrize(("cluster_size", "mailbox_depth"), [(4, 16), (8, 36)]) +def test_k1_parallel_rejects_unsafe_mailbox_depth( + flash_kda_device, monkeypatch, cluster_size, mailbox_depth +): + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128_k1_parallel", cluster_size, mailbox_depth), + ) + inputs = _make_inputs(seq_lens=[2048], num_heads=8, packed=False) + + with pytest.raises(RuntimeError, match="multiple of the helper producer count"): + recurrent_kda(**_strict_prefill_kwargs(inputs)) + + +def test_k1_parallel_rejects_unsupported_head_count(flash_kda_device, monkeypatch): + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128_k1_parallel", 4, 15), + ) + inputs = _make_inputs(seq_lens=[2048], num_heads=5, packed=False) + + with pytest.raises(RuntimeError, match="requires H == 1, H == 4"): + recurrent_kda(**_strict_prefill_kwargs(inputs)) + + def test_frozen_route_passes_nondefault_stream(cuda_device, monkeypatch): monkeypatch.setattr( kda_prefill_api, "get_compute_capability", lambda device: (10, 0) @@ -856,7 +1046,7 @@ def test_frozen_prefill_h12_packed_matches_reference(flash_kda_device): ) -def test_frozen_prefill_m64_matches_reference(flash_kda_device): +def test_frozen_prefill_m64_matches_reference(flash_kda_device, monkeypatch): inputs = _make_inputs( seq_lens=[2], num_heads=64, @@ -871,6 +1061,11 @@ def test_frozen_prefill_m64_matches_reference(flash_kda_device): expected_output, expected_state = _reference(reference_inputs) output = torch.empty_like(inputs["q"]) state_identity = inputs["initial_state"] + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m64", 0, 0), + ) actual_output, actual_state = recurrent_kda( **_strict_prefill_kwargs(inputs), @@ -894,6 +1089,244 @@ def test_frozen_prefill_m64_matches_reference(flash_kda_device): ) +@pytest.mark.parametrize("num_heads", [1, 4, 8]) +def test_frozen_prefill_m64_rejects_unsupported_heads( + flash_kda_device, monkeypatch, num_heads +): + inputs = _make_inputs( + seq_lens=[2], + num_heads=num_heads, + packed=False, + initial_state=True, + seed=2100 + num_heads, + ) + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m64", 0, 0), + ) + + with pytest.raises(RuntimeError, match="specialized for fixed N=1, H=64"): + recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + +@pytest.mark.parametrize(("seq_len", "num_heads"), [(2048, 4), (2048, 8), (2048, 16)]) +def test_k1_parallel_prefill_matches_reference(flash_kda_device, seq_len, num_heads): + variant, cluster_size, _mailbox_depth = ( + kda_prefill_api._select_flash_kda_prefill_variant( + fixed_layout=True, + num_sequences=1, + num_heads=num_heads, + sequence_length=seq_len, + device=flash_kda_device, + ) + ) + assert (variant, cluster_size) == ("m128_k1_parallel", 4) + + inputs = _make_inputs( + seq_lens=[seq_len], + num_heads=num_heads, + packed=False, + initial_state=True, + seed=9000 + num_heads, + ) + reference_inputs = { + **inputs, + "initial_state": inputs["initial_state"].clone(), + } + expected_output, expected_state = _reference(reference_inputs) + output = torch.empty_like(inputs["q"]) + + actual_output, actual_state = recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=output, + output_final_state=True, + ) + + torch.testing.assert_close( + actual_output.float(), expected_output.float(), atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + actual_state.float(), expected_state.float(), atol=1e-2, rtol=1e-2 + ) + + +def test_k1_parallel_h4_matches_m128_bitwise(cuda_device, monkeypatch): + if get_compute_capability(cuda_device) not in ((10, 0), (10, 3)): + pytest.skip("FlashKDA H4 requires an SM100-family GPU") + + inputs = _make_inputs( + seq_lens=[2048], + num_heads=4, + packed=False, + initial_state=True, + seed=9040, + ) + initial_state = inputs["initial_state"].clone() + helper_output, helper_state = recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128", 0, 0), + ) + baseline_output, baseline_state = recurrent_kda( + **_strict_prefill_kwargs({**inputs, "initial_state": initial_state}), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + torch.testing.assert_close(helper_output, baseline_output, atol=0, rtol=0) + torch.testing.assert_close(helper_state, baseline_state, atol=0, rtol=0) + + +def test_m128_k1_parallel_c8_matches_m128_bitwise(cuda_device, monkeypatch): + if get_compute_capability(cuda_device) not in ((10, 0), (10, 3)): + pytest.skip("the forced M128 C8 route requires an SM100-family GPU") + + inputs = _make_inputs( + seq_lens=[2048], + num_heads=8, + packed=False, + initial_state=True, + seed=9080, + ) + initial_state = inputs["initial_state"].clone() + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128_k1_parallel", 8, 35), + ) + helper_output, helper_state = recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128", 0, 0), + ) + baseline_output, baseline_state = recurrent_kda( + **_strict_prefill_kwargs({**inputs, "initial_state": initial_state}), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + torch.testing.assert_close(helper_output, baseline_output, atol=0, rtol=0) + torch.testing.assert_close(helper_state, baseline_state, atol=0, rtol=0) + + +@pytest.mark.parametrize( + ("seq_len", "num_heads"), + [(1024, 1), (2048, 1), (4096, 1), (1024, 4), (2048, 4)], +) +def test_m64_k1_parallel_matches_reference( + flash_kda_device, monkeypatch, seq_len, num_heads +): + inputs = _make_inputs( + seq_lens=[seq_len], + num_heads=num_heads, + packed=False, + initial_state=True, + seed=9064 + seq_len, + ) + reference_inputs = { + **inputs, + "initial_state": inputs["initial_state"].clone(), + } + expected_output, expected_state = _reference(reference_inputs) + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m64_k1_parallel", 4, 10), + ) + + actual_output, actual_state = recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + torch.testing.assert_close( + actual_output.float(), expected_output.float(), atol=1e-2, rtol=1e-2 + ) + torch.testing.assert_close( + actual_state.float(), expected_state.float(), atol=1e-2, rtol=1e-2 + ) + + +def test_m64_k1_parallel_rejects_unvalidated_c8(flash_kda_device, monkeypatch): + inputs = _make_inputs( + seq_lens=[1024], + num_heads=1, + packed=False, + initial_state=True, + seed=9164, + ) + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m64_k1_parallel", 8, 30), + ) + + with pytest.raises(RuntimeError, match="requires cluster_size == 4"): + recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + ) + + +@pytest.mark.parametrize("num_heads", [4, 8]) +def test_k1_parallel_varlen_matches_m128_bitwise(cuda_device, monkeypatch, num_heads): + if get_compute_capability(cuda_device) != (10, 0): + pytest.skip("packed-varlen K1 helpers are currently enabled only on B200") + + seq_lens = [2304, 1792] + inputs = _make_inputs( + seq_lens=seq_lens, + num_heads=num_heads, + packed=True, + initial_state=True, + seed=9017 + num_heads, + ) + initial_state = inputs["initial_state"].clone() + seq_order = torch.tensor([0, 1], dtype=torch.int32, device=cuda_device) + + helper_output, helper_state = recurrent_kda( + **_strict_prefill_kwargs(inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + seq_order=seq_order, + ) + + monkeypatch.setattr( + kda_prefill_api, + "_select_flash_kda_prefill_variant", + lambda **_kwargs: ("m128", 0, 0), + ) + baseline_inputs = {**inputs, "initial_state": initial_state} + baseline_output, baseline_state = recurrent_kda( + **_strict_prefill_kwargs(baseline_inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + seq_order=seq_order, + ) + + torch.testing.assert_close(helper_output, baseline_output, atol=0, rtol=0) + torch.testing.assert_close(helper_state, baseline_state, atol=0, rtol=0) + + @pytest.mark.parametrize( ("packed", "num_heads", "has_initial_state"), [(False, 64, True), (True, 2, False)], @@ -978,6 +1411,64 @@ def test_frozen_prefill_cuda_graph_capture_and_replay( ) +@pytest.mark.parametrize(("seq_len", "num_heads"), [(2048, 8), (4096, 1)]) +def test_k1_parallel_cuda_graph_capture_and_replay( + flash_kda_device, seq_len, num_heads +): + if num_heads == 1 and get_compute_capability(flash_kda_device) != (10, 0): + pytest.skip("the M64 owner/helper route is enabled only on B200/GB200") + + inputs = _make_inputs( + seq_lens=[seq_len], + num_heads=num_heads, + packed=False, + initial_state=True, + seed=9200 + num_heads, + ) + state_seed = inputs["initial_state"].clone() + expected_inputs = {**inputs, "initial_state": state_seed.clone()} + expected_output, expected_state = recurrent_kda( + **_strict_prefill_kwargs(expected_inputs), + output=torch.empty_like(inputs["q"]), + output_final_state=True, + prefill_workspace=RecurrentKDAPrefillWorkspace(flash_kda_device), + ) + + inputs["initial_state"].copy_(state_seed) + output = torch.empty_like(inputs["q"]) + workspace = RecurrentKDAPrefillWorkspace(flash_kda_device) + capture_stream = torch.cuda.Stream(device=flash_kda_device) + capture_stream.wait_stream(torch.cuda.current_stream(flash_kda_device)) + call_kwargs = { + **_strict_prefill_kwargs(inputs), + "output": output, + "output_final_state": True, + "prefill_workspace": workspace, + } + + with torch.cuda.stream(capture_stream): + recurrent_kda(**call_kwargs) + inputs["initial_state"].copy_(state_seed) + output.zero_() + capture_stream.synchronize() + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=capture_stream): + captured_output, captured_state = recurrent_kda(**call_kwargs) + + with torch.cuda.stream(capture_stream): + inputs["initial_state"].copy_(state_seed) + output.fill_(float("nan")) + capture_stream.synchronize() + graph.replay() + torch.cuda.synchronize() + + assert captured_output.data_ptr() == output.data_ptr() + assert captured_state is inputs["initial_state"] + torch.testing.assert_close(captured_output, expected_output, atol=0, rtol=0) + torch.testing.assert_close(captured_state, expected_state, atol=0, rtol=0) + + @pytest.mark.parametrize("num_heads", [6, 12]) def test_frozen_prefill_non_aligned_heads_graph_refreshes_beta( flash_kda_device, num_heads