From c06165092990e94e26d3cc57f780d07634bd022f Mon Sep 17 00:00:00 2001 From: Cong Zhang Date: Wed, 7 Oct 2026 04:23:45 -0700 Subject: [PATCH 1/7] [cl] Add SM100a cuda-lang ts paged FMHA decode kernel --- .../test/test_fmha_decode_sm100a.py | 305 +++++++++++ .../fmha_decode_config.py | 90 ++++ .../fmha_decode_kernel.py | 339 +++++++++++++ .../fmha_decode_resources/helpers_common.py | 117 +++++ .../fmha_decode_resources/smem_resources.py | 114 +++++ .../fmha_decode_resources/tmem_resources.py | 422 ++++++++++++++++ .../08_fmha_decode_sm100a/fmha_decode_run.py | 135 +++++ .../fmha_decode_tasks.py | 478 ++++++++++++++++++ 8 files changed, 2000 insertions(+) create mode 100644 experimental/task_scheduling/test/test_fmha_decode_sm100a.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_config.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_kernel.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/helpers_common.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/smem_resources.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/tmem_resources.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_run.py create mode 100644 experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_tasks.py diff --git a/experimental/task_scheduling/test/test_fmha_decode_sm100a.py b/experimental/task_scheduling/test/test_fmha_decode_sm100a.py new file mode 100644 index 00000000..ad4ec06b --- /dev/null +++ b/experimental/task_scheduling/test/test_fmha_decode_sm100a.py @@ -0,0 +1,305 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""End-to-end paged decode tests with runtime ragged Q and KV bounds. + +Set CUDA_LANG_DECODE_EXHAUSTIVE=1 to run all 164 cases instead of the 56 CI cases. +""" + +import importlib +import os +from dataclasses import replace +from itertools import product + +import pytest +import torch +from task_scheduling_test_utils import require_blackwell_cc100 + +_PACKAGE = "experimental.task_scheduling.tutorial.08_fmha_decode_sm100a" + + +def pytest_generate_tests(metafunc): + # Each distinct configuration compiles a large task-scheduled kernel. Cover + # every parameter value in CI; retain the full cross product for longer runs. + name = metafunc.function.__name__ + dtypes = ("float16", "bfloat16") + if name == "test_paged_decode": + names = "decode_case,page,heads_q,heads_kv,persistent" + full = ( + (dtype, page, hq, hkv, persistent) + for dtype, page, (hq, hkv), persistent in product( + dtypes, (16, 32, 64, 128), ((2, 2), (8, 2), (8, 1), (32, 1)), + (False, True), + ) + ) + cases = ( + ("float16", 16, 2, 2, False), + ("bfloat16", 16, 8, 1, True), + ("bfloat16", 32, 8, 2, False), + ("float16", 32, 32, 1, True), + ("float16", 64, 8, 1, False), + ("bfloat16", 64, 2, 2, True), + ("bfloat16", 128, 32, 1, False), + ("float16", 128, 8, 2, True), + ) + elif name in ("test_paged_decode_causal_window", "test_packed_query_causal_window"): + names = "decode_case,window,persistent" + first_window = -1 if name == "test_paged_decode_causal_window" else 1 + full = product(dtypes, (first_window, 0, 31, 128, 256), (False, True)) + cases = ( + ("float16", first_window, False), + ("float16", 0, False), + ("bfloat16", 0, True), + ("float16", 31, True), + ("bfloat16", 128, False), + ("bfloat16", 256, True), + ) + elif name == "test_packed_query_decode": + names = "decode_case,page,ratio,grouped,persistent,mask" + layouts = ( + (16, 1, True, False), (32, 2, True, True), + (64, 4, True, False), (128, 8, True, True), + (16, 3, False, False), (32, 10, False, True), + (64, 32, False, False), (128, 4, False, True), + ) + full = ( + (dtype, *layout, mask) + for dtype, layout, mask in product(dtypes, layouts, ("dense", "causal")) + ) + cases = ( + ("float16", 16, 1, True, False, "dense"), + ("bfloat16", 32, 2, True, True, "causal"), + ("bfloat16", 64, 4, True, False, "dense"), + ("float16", 128, 8, True, True, "causal"), + ("bfloat16", 16, 3, False, False, "causal"), + ("float16", 32, 10, False, True, "dense"), + ("float16", 64, 32, False, False, "causal"), + ("bfloat16", 128, 4, False, True, "dense"), + ) + else: + if "decode_case" in metafunc.fixturenames: + metafunc.parametrize("decode_case", dtypes, indirect=True, ids=("fp16", "bf16")) + return + if os.environ.get("CUDA_LANG_DECODE_EXHAUSTIVE") == "1": + cases = tuple(full) + ids = [ + "-".join(("fp16" if case[0] == "float16" else "bf16", *(str(v) for v in case[1:]))) + for case in cases + ] + metafunc.parametrize(names, cases, indirect=["decode_case"], ids=ids) + + +@pytest.fixture +def decode_case(request): + runner = importlib.import_module(_PACKAGE + ".fmha_decode_run") + dtype = getattr(runner.cl, request.param) + cfg = runner.FmhaDecodeConfig(q_dtype=dtype, kv_dtype=dtype, out_dtype=dtype) + return runner, cfg + + +@require_blackwell_cc100() +def test_paged_decode(page, heads_q, heads_kv, persistent, decode_case): + runner, cfg = decode_case + cfg = replace(cfg, num_tokens_per_page=page, use_persistent_scheduler=persistent) + tensors = runner.prepare_tensors( + (1, 17, 127, 128, 129, 255, 256, 257, 1023), heads_q, heads_kv, cfg + ) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@require_blackwell_cc100() +def test_paged_decode_causal_window(window, persistent, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, mask_type="causal", window_left=window, use_persistent_scheduler=persistent + ) + tensors = runner.prepare_tensors((1, 129, 257, 8192), 8, 1, cfg) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@require_blackwell_cc100() +def test_paged_decode_runtime_metadata(decode_case): + runner, cfg = decode_case + tensors = runner.prepare_tensors((8192,) * 4, cfg=cfg) + runner.run(tensors, cfg) + tensors["seq_lens"].copy_(torch.tensor((1, 129, 1025, 8191), device="cuda", dtype=torch.int32)) + tensors["paged_kv_indptr"].copy_( + torch.tensor((0, 1, 6, 39, 295), device="cuda", dtype=torch.int32) + ) + tensors["paged_kv_indices"].copy_(tensors["paged_kv_indices"].flip(0)) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@pytest.mark.parametrize("head_ratio", (3, 10)) +@require_blackwell_cc100() +def test_paged_decode_partial_head_group(head_ratio, decode_case): + runner, cfg = decode_case + tensors = runner.prepare_tensors((129, 1024), head_ratio * 2, 2, cfg) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@pytest.mark.parametrize("grouped", (False, True)) +@require_blackwell_cc100() +def test_paged_decode_clc_recycles_requests(grouped, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, mask_type="causal", use_persistent_scheduler=True, groups_tokens_heads_q=grouped + ) + # More than three waves, with different request lengths on recycled CTAs. + lengths = tuple((1, 129, 1025, 8192)[i % 4] for i in range(64)) + tensors = runner.prepare_tensors(lengths, 32, 8, cfg) + runner.run(tensors, cfg) + torch.cuda.synchronize() + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + runner.run(tensors, cfg) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + # Replay the same compiled graph with different CSR boundaries and page IDs. + lengths = torch.tensor( + tuple((257, 1, 513, 33)[i % 4] for i in range(64)), device="cuda", dtype=torch.int32 + ) + tensors["seq_lens"].copy_(lengths) + tensors["paged_kv_indptr"][1:].copy_(((lengths + 31) // 32).cumsum(0)) + tensors["paged_kv_indices"].copy_(tensors["paged_kv_indices"].flip(0)) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@pytest.mark.parametrize("head_ratio", (1, 2, 8)) +@require_blackwell_cc100() +def test_paged_decode_grouped_q(head_ratio, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, use_persistent_scheduler=True, groups_tokens_heads_q=True + ) + tensors = runner.prepare_tensors((1, 129, 1025), head_ratio * 2, 2, cfg) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@require_blackwell_cc100() +def test_packed_query_decode(page, ratio, grouped, persistent, mask, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, use_variable_seqlens_q=True, max_seq_len_q=17, num_tokens_per_page=page, + groups_tokens_heads_q=grouped, use_persistent_scheduler=persistent, mask_type=mask, + ) + tensors = runner.prepare_tensors( + (1, 129, 255, 256, 257, 513), ratio * 2, 2, cfg, q_lengths=(1, 3, 7, 8, 9, 17) + ) + tensors["o"].fill_(float("nan")) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@require_blackwell_cc100() +def test_packed_query_causal_window(window, persistent, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, use_variable_seqlens_q=True, max_seq_len_q=17, groups_tokens_heads_q=True, + use_persistent_scheduler=persistent, mask_type="causal", window_left=window, + ) + tensors = runner.prepare_tensors( + (1, 129, 257, 8192), 4, 2, cfg, q_lengths=(1, 3, 9, 17) + ) + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@pytest.mark.parametrize("persistent", (False, True)) +@pytest.mark.parametrize("grouped", (False, True)) +@require_blackwell_cc100() +def test_packed_query_graph_reloads_metadata(grouped, persistent, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, use_variable_seqlens_q=True, max_seq_len_q=9, groups_tokens_heads_q=grouped, + use_persistent_scheduler=persistent, mask_type="causal", window_left=31, + ) + # Active and inactive tiles cross multiple resident waves and rotate on replay. + lengths = tuple((33, 129, 257, 513)[i % 4] for i in range(24)) + q_lengths = tuple((1, 3, 7, 9)[i % 4] for i in range(24)) + ratio = 4 if grouped else 10 + tensors = runner.prepare_tensors(lengths, ratio * 4, 4, cfg, q_lengths=q_lengths) + runner.run(tensors, cfg) + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + runner.run(tensors, cfg) + for replay in range(2): + if replay: + new_q_lengths = torch.tensor(q_lengths[::-1], device="cuda", dtype=torch.int32) + tensors["qo_indptr"][1:].copy_(new_q_lengths.cumsum(0)) + new_lengths = torch.tensor(lengths[::-1], device="cuda", dtype=torch.int32) + tensors["seq_lens"].copy_(new_lengths) + tensors["paged_kv_indptr"][1:].copy_(((new_lengths + 31) // 32).cumsum(0)) + tensors["paged_kv_indices"].copy_(tensors["paged_kv_indices"].flip(0)) + tensors["o"].fill_(float("nan")) + graph.replay() + torch.cuda.synchronize() + torch.testing.assert_close( + tensors["o"], runner.torch_reference(tensors, cfg), atol=0.003, rtol=0.02 + ) + + +@pytest.mark.parametrize("persistent", (False, True)) +@require_blackwell_cc100() +def test_packed_query_window_zero_selects_exact_key(persistent, decode_case): + runner, cfg = decode_case + cfg = replace( + cfg, use_variable_seqlens_q=True, max_seq_len_q=9, groups_tokens_heads_q=True, + use_persistent_scheduler=persistent, mask_type="causal", window_left=0, + ) + lengths, q_lengths = (17, 129, 257, 33), (1, 7, 3, 9) + tensors = runner.prepare_tensors(lengths, 4, 2, cfg, q_lengths=q_lengths) + tensors["q"].zero_() + tensors["k"].zero_() + tensors["v"].fill_(float("nan")) + expected = torch.empty_like(tensors["o"]) + q_begin = 0 + for batch, (length, q_length) in enumerate(zip(lengths, q_lengths)): + begin, end = tensors["paged_kv_indptr"][batch:batch + 2].tolist() + ids = tensors["paged_kv_indices"][begin:end].long() + # Unique keys make an off-by-one query/causal offset observable exactly. + values = torch.arange((end - begin) * 32, device="cuda", dtype=torch.float32) + values = values + batch * 16 + tensors["v"][ids] = values.reshape(-1, 1, 32, 1).to(tensors["v"].dtype) + expected[q_begin:q_begin + q_length] = values[length - q_length:length, None, None] + q_begin += q_length + runner.run(tensors, cfg) + torch.cuda.synchronize() + torch.testing.assert_close(tensors["o"], expected, atol=0, rtol=0) diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_config.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_config.py new file mode 100644 index 00000000..d9173c22 --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_config.py @@ -0,0 +1,90 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Static configuration for the D128, Q8 SwapsMmaAb decode port.""" + +from dataclasses import dataclass, fields + +import cuda.lang as cl + +from cuda.tile import DType + + +@dataclass(frozen=True) +class FmhaDecodeConfig: + headdim: int = 128 + tile_size_q: int = 8 + tile_size_kv: int = 128 + num_insts_kv: int = 2 + q_dtype: DType = cl.float16 + kv_dtype: DType = cl.float16 + out_dtype: DType = cl.float16 + q_stages: int = 2 + kv_stages: int = 4 + o_stages: int = 2 + page_offsets_stages: int = 6 + num_tokens_per_page: int = 32 + softmax0_warp_idx: int = 0 + softmax1_warp_idx: int = 4 + correction_warp_idx: int = 8 + mma_warp_idx: int = 12 + load_warp_idx: int = 13 + scheduler_warp_idx: int = 13 + page_offsets_warp_idx: int = 14 + clc_load_warp_idx: int = 15 + threads_per_cta: int = 512 + tmem_s_cols: int = 8 + tmem_stats_cols: int = 32 + tmem_o_cols: int = 8 + tmem_alloc_cols: int = 128 + window_left: int = -1 + mask_type: str = "dense" + use_persistent_scheduler: bool = False + groups_tokens_heads_q: bool = False + use_variable_seqlens_q: bool = False + max_seq_len_q: int = 1 + heads_q_per_kv: int = 0 + + @property + def q_tokens_per_cta(self): + return self.tile_size_q // self.heads_q_per_kv if self.groups_tokens_heads_q else 1 + + def num_q_ctas(self, head_ratio): + if self.groups_tokens_heads_q: + tokens = self.tile_size_q // head_ratio + return (self.max_seq_len_q + tokens - 1) // tokens + return self.max_seq_len_q * ((head_ratio + self.tile_size_q - 1) // self.tile_size_q) + + +def validate_config(cfg): + defaults = FmhaDecodeConfig() + for field in fields(cfg): + if field.name not in ( + "num_tokens_per_page", "mask_type", "window_left", + "use_persistent_scheduler", "groups_tokens_heads_q", + "q_dtype", "kv_dtype", "out_dtype", + "use_variable_seqlens_q", "max_seq_len_q", "heads_q_per_kv", + ): + if getattr(cfg, field.name) != getattr(defaults, field.name): + raise ValueError( + f"{field.name} is fixed at {getattr(defaults, field.name)} " + "in this specialization" + ) + if (cfg.headdim, cfg.tile_size_q, cfg.tile_size_kv, cfg.num_insts_kv) != (128, 8, 128, 2): + raise ValueError("This port currently implements D128/Q8/KV128 with two K/V instances") + if cfg.q_dtype not in (cl.float16, cl.bfloat16): + raise ValueError("Q/K/V/O dtype must be float16 or bfloat16") + if not cfg.q_dtype == cfg.kv_dtype == cfg.out_dtype: + raise ValueError("Q, K, V, and O must use the same dtype") + if cfg.num_tokens_per_page not in (16, 32, 64, 128): + raise ValueError("page size must be 16, 32, 64, or 128") + if cfg.mask_type not in ("dense", "causal"): + raise ValueError("mask_type must be dense or causal") + if cfg.window_left < -1 or (cfg.window_left >= 0 and cfg.mask_type != "causal"): + raise ValueError("window_left requires causal attention and must be >= -1") + if not isinstance(cfg.max_seq_len_q, int) or cfg.max_seq_len_q < 1: + raise ValueError("max_seq_len_q must be a positive integer") + if not cfg.use_variable_seqlens_q and cfg.max_seq_len_q != 1: + raise ValueError("multiple query tokens currently require packed Q and qo_indptr") + if not 0 <= cfg.heads_q_per_kv <= 32: + raise ValueError("heads_q_per_kv must be 0 (inferred) or an integer in [1,32]") diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_kernel.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_kernel.py new file mode 100644 index 00000000..eccc251a --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_kernel.py @@ -0,0 +1,339 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""D128 paged decode with fixed SQ1 or packed queries.""" + +from dataclasses import dataclass, replace +from functools import lru_cache + +import cuda.lang as cl +import task_scheduling as ts + +from .fmha_decode_config import FmhaDecodeConfig, validate_config +from .fmha_decode_resources.helpers_common import q_group_token_base +from .fmha_decode_resources.smem_resources import ( + SmemQResource, + SmemKvResource, + SmemPageOffsetsKvResource, +) +from .fmha_decode_resources.tmem_resources import ( + TmemSResource, + SmemPResource, + TmemSoftmaxStatsResource, + TmemOResource, +) +from .fmha_decode_tasks import ( + create_load_task, + create_page_offsets_task, + create_mma_task, + create_softmax_task, + create_correction_task, + create_scheduler_task, + ScheduleTokenThrottleResource, + PackedDecodeWorkQueue, +) + + +@dataclass(frozen=True) +class TasksInputs: + tma_q_desc: object + tma_k_desc: object + tma_v_desc: object + o: object + seq_len: object + page_begin: object + page_count: object + paged_kv_indices: object + head_ratio: object + tmem_base: object + seq_lens: object + paged_kv_indptr: object + qo_indptr: object + q_token_offset: object + seq_len_q: object + + +def _pipeline(kind, stages, producers, consumers, **kwargs): + method = getattr(ts.PipelineConfig, "create_" + kind + "_pipeline_cfg") + return method( + num_stages=stages, + producer_group=ts.CooperativeGroup(producers), + consumer_group=ts.CooperativeGroup(consumers), + cta_layout_vmnk=(1, 1, 1, 1), + producer_signaling_threads=( + ts.SignalingThreads.TaskWarpLeader + if kind == "tma_umma" + else ts.SignalingThreads.CtaLeader if kind == "umma_async" else ts.SignalingThreads.All + ), + consumer_signaling_threads=( + ts.SignalingThreads.CtaLeader if kind.endswith("umma") else ts.SignalingThreads.All + ), + advance_on_wait=True, + tcgen05_fence_after_wait=False, + **kwargs, + ) + + +def build_fmha_decode_task_manager( + cfg=FmhaDecodeConfig(), *, problem_shape=(1, 1, 1), verbose=False +): + validate_config(cfg) + smem = ts.SmemAllocator(default_add_barriers=False) + allocations = {} + for name, size, alignment in ( + ("q", 4096, 1024), + ("kv", 131072, 1024), + ("p0", 2048, 1024), + ("p1", 2048, 1024), + ("output", 2048, 1024), + ("pages", 768, 128), + ("max0", 32, 16), + ("max1", 32, 16), + ("sum", 128, 16), + ("tmem_ptr", 4, 4), + ): + allocations[name] = ts.SmemAllocation(name, size, alignment=alignment) + smem.add(allocations[name]) + work_queue = None + throttle = None + if cfg.use_persistent_scheduler: + response = ts.SmemAllocation("clc_response", 32, alignment=16, count=2) + smem.add(response) + queue_cls = PackedDecodeWorkQueue if cfg.use_variable_seqlens_q else ts.WorkQueue + work_queue = queue_cls( + **({"cfg": cfg} if cfg.use_variable_seqlens_q else {}), + name="work_queue", + pipeline_config=ts.PipelineConfig.create_clc_fetch_async_pipeline_cfg( + num_stages=2, + num_bytes=16, + producer_group=ts.CooperativeGroup(1), + consumer_group=ts.CooperativeGroup(512), + cta_layout_vmnk=(1, 1, 1, 1), + producer_signaling_threads=ts.SignalingThreads.CtaLeader, + consumer_signaling_threads=ts.SignalingThreads.All, + ), + tile_scheduler_config=( + ts.TileSchedulerConfig.create_clc_dynamic_persistent_tile_scheduler_params( + ts.ClcDynamicPersistentTileSchedulerParams(problem_shape, (1, 1, 1)), + response, + ) + ), + ) + throttle = ScheduleTokenThrottleResource( + name="schedule_token_throttle", + pipeline_config=_pipeline("async_async", 2, 32, 32), + ) + smem.compute_layout() + offsets = {name: allocation.offset for name, allocation in allocations.items()} + sq = SmemQResource( + name="smem_q", pipeline_config=_pipeline("tma_umma", 2, 1, 1, num_bytes=2048) + ) + skv = SmemKvResource( + name="smem_kv", pipeline_config=_pipeline("tma_umma", 4, 1, 1, num_bytes=32768) + ) + pages = SmemPageOffsetsKvResource( + name="smem_page_offsets_kv", pipeline_config=_pipeline("async_async", 6, 32, 32) + ) + s0 = TmemSResource(name="tmem_s0", pipeline_config=_pipeline("umma_async", 1, 1, 128)) + s1 = TmemSResource(name="tmem_s1", pipeline_config=_pipeline("umma_async", 1, 1, 128)) + p0 = SmemPResource(name="smem_p0", pipeline_config=_pipeline("async_umma", 1, 128, 1)) + p1 = SmemPResource(name="smem_p1", pipeline_config=_pipeline("async_umma", 1, 128, 1)) + stats0 = TmemSoftmaxStatsResource( + name="tmem_stats0", pipeline_config=_pipeline("async_async", 1, 128, 128) + ) + stats1 = TmemSoftmaxStatsResource( + name="tmem_stats1", pipeline_config=_pipeline("async_async", 1, 128, 128) + ) + o = TmemOResource(name="tmem_o", pipeline_config=_pipeline("umma_async", 2, 1, 128)) + resources = (sq, skv, pages, s0, s1, p0, p1, stats0, stats1, o) + if work_queue is not None: + resources += (work_queue, throttle) + barriers = ts.BarrierAllocator() + for resource in resources: + barriers.add_resource(resource) + barriers.compute_layout() + tmem = ts.TmemAllocator() + for name, columns in (("s0", 8), ("s1", 8), ("stats0", 32), ("stats1", 32), ("o", 16)): + tmem.add(ts.TmemAllocation(name, columns)) + tmem.compute_layout() + tasks = [ + create_softmax_task(s0, p0, stats0, 0, offsets, cfg, work_queue), + create_softmax_task(s1, p1, stats1, 1, offsets, cfg, work_queue), + create_correction_task(stats0, stats1, o, offsets, cfg, work_queue), + create_mma_task(sq, skv, s0, s1, p0, p1, o, offsets, cfg, work_queue), + create_load_task(sq, skv, pages, offsets, cfg, work_queue, throttle), + create_page_offsets_task(pages, offsets["pages"], cfg, work_queue), + ] + dependencies = { + sq: [], + pages: [], + skv: [pages], + s0: [sq, skv], + s1: [sq, skv], + p0: [s0], + p1: [s1], + stats0: [s0], + stats1: [s1], + o: [skv, p0, p1], + } + if work_queue is not None: + tasks.append(create_scheduler_task(work_queue, throttle, cfg)) + for deps in dependencies.values(): + deps.append(work_queue) + dependencies[work_queue] = [work_queue, throttle] + dependencies[throttle] = [work_queue] + manager = ts.TaskManager( + tasks=tasks, + resource_dependency_graph=dependencies, + smem_allocator=smem, + tmem_allocator=tmem, + barrier_allocator=barriers, + cta_warps=16, + verbose=verbose, + exhaustive_deadlock_race_check=False, + ) + return manager + + +@lru_cache(maxsize=None) +def make_fmha_decode_kernel(head_ratio, heads_kv, cfg=FmhaDecodeConfig(), *, batch=1): + if cfg.heads_q_per_kv not in (0, head_ratio): + raise ValueError("heads_q_per_kv does not match the input tensors") + if cfg.groups_tokens_heads_q and head_ratio not in (1, 2, 4, 8): + raise ValueError("grouped Q8 supports Hq/Hkv = 1, 2, 4, or 8") + cfg = replace(cfg, heads_q_per_kv=head_ratio) + device_manager = build_fmha_decode_task_manager( + cfg, problem_shape=(cfg.num_q_ctas(head_ratio), heads_kv, batch) + ).to_device() + + def fmha_decode_impl(tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr): + if cfg.use_variable_seqlens_q: + seq_lens = cl.Array.from_parts(seq_lens.pointer(), seq_lens.shape, (1,)) + paged_kv_indptr = cl.Array.from_parts( + paged_kv_indptr.pointer(), paged_kv_indptr.shape, (1,) + ) + paged_kv_indices = cl.Array.from_parts( + paged_kv_indices.pointer(), paged_kv_indices.shape, (1,) + ) + qo_indptr = cl.Array.from_parts(qo_indptr.pointer(), qo_indptr.shape, (1,)) + row_stride = head_ratio * heads_kv * 128 + o = cl.Array.from_parts( + o.pointer(), (o.shape[0], head_ratio * heads_kv, 128), (row_stride, 128, 1) + ) + q_token_offset, seq_len_q = cl.int32(0), cl.int32(1) + q_tile_is_active = True + if cfg.use_variable_seqlens_q and not cfg.use_persistent_scheduler: + request = cl.block_index(2) + q_token_offset = qo_indptr[request] + seq_len_q = qo_indptr[request + 1] - q_token_offset + q_tile_is_active = q_group_token_base(cl.block_index(0), cfg) < seq_len_q + allocators = device_manager.setup_resources_and_tasks() + if q_tile_is_active: + ptr = allocators.smem_allocator.get( + "tmem_ptr", cl.pointer_dtype(cl.float32, cl.MemorySpace.TENSOR) + ) + warp, _ = cl.shuffle_sync(cl.ShuffleKind.INDEX, cl.thread_index(0) // 32, 0) + if warp == 12: + cl.tcgen05_allocate(ptr.pointer(), 128, cta_group=cl.CTAGroup.CTA_1) + cl.tcgen05_relinquish_allocation_permit(cta_group=cl.CTAGroup.CTA_1) + cl.barrier_sync_block_aligned() + base = ptr[0] + if cfg.use_persistent_scheduler: + # Each worker binds these fields at the head of its work-tile loop. + seq_len = cl.int32(0) + page_begin = cl.int32(0) + page_count = cl.int32(0) + else: + request = cl.block_index(2) + seq_len = seq_lens[request] + page_begin = paged_kv_indptr[request] + page_count = paged_kv_indptr[request + 1] - page_begin + inputs = TasksInputs( + tq, tk, tv, o, seq_len, page_begin, page_count, paged_kv_indices, head_ratio, base, + seq_lens, paged_kv_indptr, + qo_indptr, q_token_offset, seq_len_q, + ) + device_manager.run(inputs, allocators) + cl.barrier_sync_block_aligned() + if warp == 12: + cl.tcgen05_deallocate(base, 128, cta_group=cl.CTAGroup.CTA_1) + + if cfg.use_variable_seqlens_q: + @cl.kernel(max_threads_per_block=(512,), min_blocks_per_sm=1) + def fmha_decode(tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr): + fmha_decode_impl(tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr) + else: + @cl.kernel(max_threads_per_block=(512,), min_blocks_per_sm=1) + def fmha_decode(tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices): + fmha_decode_impl(tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices, None) + + return fmha_decode + + +@lru_cache(maxsize=None) +def make_fmha_decode_launcher(head_ratio, heads_kv, cfg=FmhaDecodeConfig(), *, batch=1): + """Encode the Q/K/V tensor maps on the host before launching decode.""" + kernel = make_fmha_decode_kernel(head_ratio, heads_kv, cfg, batch=batch) + cfg = replace(cfg, heads_q_per_kv=head_ratio) + grid = (cfg.num_q_ctas(head_ratio), heads_kv, batch) + q_box = (64, 8, 1, 1, 1) + if cfg.groups_tokens_heads_q: + q_box = (64, head_ratio, 1, 8 // head_ratio, 1) + if cfg.use_variable_seqlens_q: + q_box = (64, head_ratio if cfg.groups_tokens_heads_q else 8, cfg.q_tokens_per_cta, 1, 1) + + def launch_impl(stream, q, k, v, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr): + row_stride = head_ratio * heads_kv * 128 + if cfg.use_variable_seqlens_q: + q_view = cl.Array.from_parts( + q.pointer(), + (128, head_ratio * heads_kv, cfg.q_tokens_per_cta, 1 << 31, 1 << 31), + (1, 128, row_stride, (1 << 35) - row_stride, row_stride), + ) + else: + q_view = cl.Array.from_parts( + q.pointer(), + (128, head_ratio, heads_kv, 1, q.shape[0]), + (1, 128, head_ratio * 128, row_stride, row_stride), + ) + kv_stride = cfg.num_tokens_per_page * 128 + k_view = cl.Array.from_parts( + k.pointer(), + (128, cfg.num_tokens_per_page, heads_kv, k.shape[0]), + (1, 128, kv_stride, heads_kv * kv_stride), + ) + v_view = cl.Array.from_parts( + v.pointer(), + (128, cfg.num_tokens_per_page, heads_kv, v.shape[0]), + (1, 128, kv_stride, heads_kv * kv_stride), + ) + tq = cl.tensor_map_tiled( + q_view, q_box, order=(0, 1, 2, 3, 4), + swizzle=cl.SwizzleMode.SWIZZLE_128B, + l2_promotion=cl.TensorMapL2Promotion.NONE, + ) + tk = cl.tensor_map_tiled( + k_view, (64, cfg.num_tokens_per_page, 1, 1), order=(0, 1, 2, 3), + swizzle=cl.SwizzleMode.SWIZZLE_128B, + l2_promotion=cl.TensorMapL2Promotion.NONE, + ) + tv = cl.tensor_map_tiled( + v_view, (64, cfg.num_tokens_per_page, 1, 1), order=(0, 1, 2, 3), + swizzle=cl.SwizzleMode.SWIZZLE_128B, + l2_promotion=cl.TensorMapL2Promotion.NONE, + ) + args = (tq, tk, tv, o, seq_lens, paged_kv_indptr, paged_kv_indices) + if cfg.use_variable_seqlens_q: + args += (qo_indptr,) + cl.launch(stream, grid, (512,), kernel, args, programmatic_dependent_launch=False) + + if cfg.use_variable_seqlens_q: + @cl.host_entry + def launcher(stream, q, k, v, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr): + launch_impl(stream, q, k, v, o, seq_lens, paged_kv_indptr, paged_kv_indices, qo_indptr) + else: + @cl.host_entry + def launcher(stream, q, k, v, o, seq_lens, paged_kv_indptr, paged_kv_indices): + launch_impl(stream, q, k, v, o, seq_lens, paged_kv_indptr, paged_kv_indices, None) + + return launcher diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/helpers_common.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/helpers_common.py new file mode 100644 index 00000000..abae2753 --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/helpers_common.py @@ -0,0 +1,117 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Device address helpers shared by decode resources.""" + +import cuda.lang as cl + +NEG_MAX = -3.4028234663852886e38 + + +def smem_array(stage_info, offset, dtype, count): + ptr = stage_info.context.smem_base.pointer() + offset + ptr = cl.bitcast(ptr, cl.pointer_dtype(dtype, ptr.memory_space)) + return cl.Array.from_parts(ptr, (count,)) + + +def tmem_ptr(stage_info, column, row=0): + return cl.tcgen05_tmem_offset( + stage_info.context.tasks_inputs.tmem_base, + column_offset=column, + lane_offset=row, + ) + + +def smem_descriptor(ptr, leading): + return cl.int64( + cl.Tcgen05SharedMemoryDescriptor( + matrix_start_address=ptr, + leading_dimension_byte_offset=leading, + stride_dimension_byte_offset=1024, + swizzle_mode=cl.SwizzleMode.SWIZZLE_128B, + ).encode() + ) + + +def warp_and_lane(stage_info): + return stage_info.context.warp_index % 4, cl.thread_index(0) % 32 + + +def work_coords(stage_info): + tile = stage_info.work_tile.tile_idx + return tile[0], tile[1], tile[2] + + +def q_group_token_base(group, cfg): + if cfg.groups_tokens_heads_q: + return group * cfg.q_tokens_per_cta + return group // ((cfg.heads_q_per_kv + 7) // 8) + + +def q_row_token_and_local_head(group, row, cfg): + if cfg.groups_tokens_heads_q: + return group * cfg.q_tokens_per_cta + row // cfg.heads_q_per_kv, row % cfg.heads_q_per_kv + head_ctas = (cfg.heads_q_per_kv + 7) // 8 + return group // head_ctas, (group % head_ctas) * 8 + row + + +def kv_tile_bounds(length, seq_len_q, group, cfg): + """Union of the causal/window spans for the live Q rows in one CTA.""" + first = cl.int32(0) + last = length + if cfg.use_variable_seqlens_q: + q_begin = q_group_token_base(group, cfg) + if cfg.mask_type == "causal": + q_end = cl.minimum(q_begin + cfg.q_tokens_per_cta, seq_len_q) + last = cl.minimum(cl.maximum(length - seq_len_q + q_end, 0), length) + if cfg.window_left >= 0: + first = cl.maximum(length - seq_len_q + q_begin - cfg.window_left, 0) >> 7 + elif cfg.window_left >= 0: + first = cl.maximum(length - cfg.window_left - 1, 0) // 128 + if cfg.use_variable_seqlens_q: + return first, cl.int32(cl.uint32(last + 127) >> 7) + return first, (last + 127) // 128 + + +def transform_ragged_coords(coords, box, extent): + """Bound a token box using two synthetic TMA dimensions.""" + extent = cl.minimum(cl.maximum(extent, 0), box) + shift = (box - extent) % box + box * cl.int32(extent == 0) + return coords[0], coords[1], shift, cl.int32(1 << 30), coords[2] + (1 << 30) - shift + + +def kv_tile_idx(stage_info, inst_idx, section, is_v, cfg): + if section == 0: + tile = cl.int32(inst_idx) + elif section == 1: + tile = 2 * stage_info.loop_offset + inst_idx + (0 if is_v else 2) + else: + tile = 2 * (num_pairs(stage_info, cfg) - 1) + inst_idx + inputs = stage_info.context.tasks_inputs + start, _ = kv_tile_bounds(inputs.seq_len, inputs.seq_len_q, work_coords(stage_info)[0], cfg) + return tile + start + + +def num_pairs(stage_info, cfg): + inputs = stage_info.context.tasks_inputs + start, end = kv_tile_bounds(inputs.seq_len, inputs.seq_len_q, work_coords(stage_info)[0], cfg) + if cfg.use_variable_seqlens_q: + return (end - start + 1) >> 1 + return (end - start + 1) // 2 + + +def load_o(stage_info, column): + warp, _ = warp_and_lane(stage_info) + chunks = tuple( + cl.tcgen05_load( + cl.Tcgen05LoadStoreShape.SHAPE_16X256B, + tmem_ptr(stage_info, column, warp * 32 + chunk * 16), + element_count=4, + dtype=cl.float32, + ) + for chunk in cl.static_iter(range(2)) + ) + cl.tcgen05_wait_load() + return cl.Vector( + *tuple(chunks[i // 4][i % 4] for i in cl.static_iter(range(8))), dtype=cl.float32 + ) diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/smem_resources.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/smem_resources.py new file mode 100644 index 00000000..5ec6db90 --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/smem_resources.py @@ -0,0 +1,114 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Q/K/V TMA staging and the page-table producer resource.""" + +from dataclasses import dataclass + +import cuda.lang as cl +import task_scheduling as ts + +from .helpers_common import ( + kv_tile_idx, q_row_token_and_local_head, smem_array, smem_descriptor, + transform_ragged_coords, work_coords, +) + + +@dataclass(kw_only=True, eq=False) +class SmemQResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def tma_load(stage_info, offset, cfg): + q_group, kv_head, batch = work_coords(stage_info) + inputs = stage_info.context.tasks_inputs + q_smem = smem_array(stage_info, offset, cfg.q_dtype, cfg.q_stages * 1024) + if cfg.use_variable_seqlens_q: + token, head = q_row_token_and_local_head(q_group, cl.int32(0), cfg) + if cl.elect_sync(): + for chunk in cl.static_iter(range(2)): + coords = (chunk * 64, q_group * 8, kv_head, 0, batch) + if cfg.use_variable_seqlens_q: + coords = transform_ragged_coords( + (chunk * 64, kv_head * cfg.heads_q_per_kv + head, + inputs.q_token_offset + token), + cfg.q_tokens_per_cta, inputs.seq_len_q - token, + ) + cl.copy_async_bulk_tensor_global_to_shared( + inputs.tma_q_desc, + coords, + q_smem.pointer(stage_info.stage_idx * 1024 + chunk * 512), + stage_info.barrier, + ) + + @ts.consumer_work(outputs=1) + @staticmethod + def q_desc(stage_info, offset, cfg): + q_smem = smem_array(stage_info, offset, cfg.q_dtype, cfg.q_stages * 1024) + return smem_descriptor(q_smem.pointer(stage_info.stage_idx * 1024), 1024) + + +@dataclass(kw_only=True, eq=False) +class SmemPageOffsetsKvResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def load_page_offsets(stage_info, offset, inst_idx, section, is_v, cfg): + inputs = stage_info.context.tasks_inputs + tile = kv_tile_idx(stage_info, inst_idx, section, is_v, cfg) + if not cfg.use_variable_seqlens_q: + tile = cl.minimum(tile, (inputs.seq_len - 1) // 128) + fragments = 128 // cfg.num_tokens_per_page + logical_page = tile * fragments + window = logical_page // 32 * 32 + if cfg.use_variable_seqlens_q: + window = (logical_page >> 5) << 5 + lane = cl.thread_index(0) % 32 + begin = inputs.page_begin + count = inputs.page_count + page_idx = cl.minimum(window + lane, count - 1) + page_id = inputs.paged_kv_indices[begin + page_idx] + cache = smem_array(stage_info, offset, cl.int32, cfg.page_offsets_stages * 32) + cache[stage_info.stage_idx * 32 + lane] = page_id + + @ts.consumer_work(outputs=1) + @staticmethod + def page_ids(stage_info, offset, inst_idx, section, is_v, cfg): + inputs = stage_info.context.tasks_inputs + tile = kv_tile_idx(stage_info, inst_idx, section, is_v, cfg) + if not cfg.use_variable_seqlens_q: + tile = cl.minimum(tile, (inputs.seq_len - 1) // 128) + fragments = 128 // cfg.num_tokens_per_page + cache = smem_array(stage_info, offset, cl.int32, cfg.page_offsets_stages * 32) + begin = stage_info.stage_idx * 32 + (tile * fragments) % 32 + if cfg.use_variable_seqlens_q: + begin = stage_info.stage_idx * 32 + ((tile * fragments) & 31) + return cl.Vector(*tuple(cache[begin + i] for i in cl.static_iter(range(fragments)))) + + +@dataclass(kw_only=True, eq=False) +class SmemKvResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def tma_load(stage_info, page_ids, offset, is_v, cfg): + _, kv_head, _ = work_coords(stage_info) + inputs = stage_info.context.tasks_inputs + desc = inputs.tma_v_desc if is_v else inputs.tma_k_desc + smem = smem_array(stage_info, offset, cfg.kv_dtype, cfg.kv_stages * 16384) + if cl.elect_sync(): + for fragment in cl.static_iter(range(128 // cfg.num_tokens_per_page)): + for chunk in cl.static_iter(range(2)): + cl.copy_async_bulk_tensor_global_to_shared( + desc, + (chunk * 64, 0, kv_head, page_ids[fragment]), + smem.pointer( + stage_info.stage_idx * 16384 + + chunk * 8192 + + fragment * cfg.num_tokens_per_page * 64 + ), + stage_info.barrier, + ) + + @ts.consumer_work(outputs=1) + @staticmethod + def kv_desc(stage_info, offset, cfg): + smem = smem_array(stage_info, offset, cfg.kv_dtype, cfg.kv_stages * 16384) + return smem_descriptor(smem.pointer(stage_info.stage_idx * 16384), 16384) diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/tmem_resources.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/tmem_resources.py new file mode 100644 index 00000000..f55998a1 --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_resources/tmem_resources.py @@ -0,0 +1,422 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Swapped-operand QK/PV, online softmax, and the two-instance epilogue.""" + +from dataclasses import dataclass + +import cuda.lang as cl +import task_scheduling as ts + +from .helpers_common import ( + NEG_MAX, + load_o, + smem_array, + smem_descriptor, + tmem_ptr, + warp_and_lane, + work_coords, + num_pairs, + kv_tile_bounds, + q_group_token_base, + q_row_token_and_local_head, +) + + +def pair_vector(value): + return cl.Vector(cl.float32(value), cl.float32(value)) + + +def max_encode(value): + bits = cl.bitcast(value, cl.int32) + return cl.bitcast(bits ^ ((bits >> 31) | cl.int32(-2147483648)), cl.uint32) + + +def max_decode(value): + bits = cl.bitcast(value, cl.int32) + return cl.bitcast(bits ^ ((~(bits >> 31)) | cl.int32(-2147483648)), cl.float32) + + +def rescale(old, new): + return ( + cl.float32(0) + if old == NEG_MAX + else cl.exp2((old - new) * (128**-0.5 * 1.4426950408889634), flush_to_zero=True) + ) + + +def safe_norm_rcp(value): + """Clamp the denominator before taking an approximate reciprocal.""" + denominator = cl._nvvm.fmax_ftz_f(value, cl.float32(1.0e-12)) + return cl._inline_ptx("rcp.approx.f32 %0, %1;", cl.float32, denominator)[0] + + +def store_o(stage_info, column, values): + warp, _ = warp_and_lane(stage_info) + for chunk in cl.static_iter(range(2)): + cl.tcgen05_store( + cl.Tcgen05LoadStoreShape.SHAPE_16X256B, + tmem_ptr(stage_info, column, warp * 32 + chunk * 16), + cl.Vector(*tuple(values[chunk * 4 + i] for i in cl.static_iter(range(4)))), + ) + cl.tcgen05_wait_store() + + +@dataclass(kw_only=True, eq=False) +class TmemSResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def qk_mma(stage_info, q_desc, k_desc, column, cfg): + instr = cl.Tcgen05InstructionDescriptor( + d_type=cl.float32, + a_type=cfg.q_dtype, + b_type=cfg.q_dtype, + m=128, + n=8, + ).encode() + if cl.elect_sync(): + for ki in cl.static_iter(range(8)): + cl.tcgen05_mma( + cl.Tcgen05MMAKind.F16, + tmem_ptr(stage_info, column), + k_desc + ki * 2 + (ki // 4) * 1016, + q_desc + ki * 2 + (ki // 4) * 56, + instr, + accumulate=ki > 0, + cta_group=cl.CTAGroup.CTA_1, + ) + + @ts.consumer_work(outputs=2, work_attrs=ts.WorkAttr.AUXILIARY) + @staticmethod + def init_state(stage_info): + return pair_vector(NEG_MAX), pair_vector(0) + + @ts.consumer_work(outputs=1, work_attrs=ts.WorkAttr.AUXILIARY) + @staticmethod + def reduce_sums(stage_info, old_max, new_max, old_sum, local_sum): + scales = cl.Vector(rescale(old_max[0], new_max[0]), rescale(old_max[1], new_max[1])) + return cl.fma(old_sum, scales, local_sum) + + @ts.consumer_work(outputs=2) + @staticmethod + def compute_softmax_loop(stage_info, old_max, column, scratch_offset, inst_idx, cfg): + warp, lane = warp_and_lane(stage_info) + group, _, _ = work_coords(stage_info) + inputs = stage_info.context.tasks_inputs + scratch = smem_array(stage_info, scratch_offset, cl.uint32, 8) + if not cfg.use_variable_seqlens_q and stage_info.loop_offset == 0: + if warp == 0 and lane < 8: + scratch[lane] = max_encode(cl.float32(NEG_MAX)) + cl.barrier_sync_block(number_of_threads=128, barrier_id=inst_idx + 1) + values = load_o(stage_info, column) + length = inputs.seq_len + if not cfg.use_variable_seqlens_q: + tile = 2 * stage_info.loop_offset + inst_idx + lower = cl.int32(0) + if cfg.window_left >= 0: + lower = cl.maximum(length - cfg.window_left - 1, 0) + tile += lower // 128 + scores = values + if ((tile + 1) * 128 > length) | (tile * 128 < lower): + token_base = tile * 128 + warp * 32 + lane // 4 + scores = cl.Vector( + *tuple( + values[i] + if (token_base + (i // 2) * 8 < length) + & (token_base + (i // 2) * 8 >= lower) + else cl.float32(NEG_MAX) + for i in cl.static_iter(range(8)) + ) + ) + if group * 8 + 8 > inputs.head_ratio: + scores = cl.Vector( + *tuple( + scores[i] + if group * 8 + (lane % 4) * 2 + i % 2 < inputs.head_ratio + else cl.float32(NEG_MAX) + for i in cl.static_iter(range(8)) + ) + ) + if cfg.use_variable_seqlens_q: + start, _ = kv_tile_bounds(length, inputs.seq_len_q, group, cfg) + tile = 2 * stage_info.loop_offset + inst_idx + start + q_base = q_group_token_base(group, cfg) + visible_end, visible_begin = length, cl.int32(0) + if cfg.mask_type == "causal": + visible_end = length - inputs.seq_len_q + q_base + 1 + if cfg.window_left >= 0: + q_end = cl.minimum(q_base + cfg.q_tokens_per_cta, inputs.seq_len_q) + visible_begin = cl.maximum( + length - inputs.seq_len_q + q_end - cfg.window_left - 1, 0 + ) + scores = values + if ((tile + 1) * 128 > visible_end) | (tile * 128 < visible_begin): + masked = () + for i in cl.static_iter(range(8)): + q_token, _ = q_row_token_and_local_head( + group, (lane % 4) * 2 + i % 2, cfg + ) + upper, lower = length, cl.int32(0) + if cfg.mask_type == "causal": + upper = length - inputs.seq_len_q + q_token + 1 + if cfg.window_left >= 0: + lower = cl.maximum(upper - cfg.window_left - 1, 0) + token = tile * 128 + warp * 32 + lane // 4 + (i // 2) * 8 + masked += ( + values[i] if (token < upper) & (token >= lower) else cl.float32(NEG_MAX), + ) + scores = cl.Vector(*masked) + valid_scores = () + if cfg.groups_tokens_heads_q: + valid_tokens = cl.minimum( + cl.maximum(inputs.seq_len_q - q_base, 0), cfg.q_tokens_per_cta + ) + valid_rows = valid_tokens * cfg.heads_q_per_kv + for i in cl.static_iter(range(8)): + row = (lane % 4) * 2 + i % 2 + valid_scores += ( + scores[i] if row < valid_rows else cl.float32(NEG_MAX), + ) + else: + for i in cl.static_iter(range(8)): + q_token, head = q_row_token_and_local_head( + group, (lane % 4) * 2 + i % 2, cfg + ) + valid_scores += ( + scores[i] if (q_token < inputs.seq_len_q) & (head < cfg.heads_q_per_kv) + else cl.float32(NEG_MAX), + ) + scores = cl.Vector(*valid_scores) + maxima = () + if cfg.use_variable_seqlens_q: + local_maxima = () + for pair in cl.static_iter(range(2)): + first = cl._nvvm.fmax_ftz_f(scores[pair], scores[pair + 2]) + second = cl._nvvm.fmax_ftz_f(scores[pair + 4], scores[pair + 6]) + value = cl._nvvm.fmax_ftz_f(first, second) + local_maxima += (cl._nvvm.fmax_ftz_f(value, old_max[pair]),) + warp_maxima = () + for pair in cl.static_iter(range(2)): + value = local_maxima[pair] + shuffled, _ = cl.shuffle_sync(cl.ShuffleKind.XOR, value, 16) + value = cl._nvvm.fmax_ftz_f(value, shuffled) + shuffled, _ = cl.shuffle_sync(cl.ShuffleKind.XOR, value, 8) + value = cl._nvvm.fmax_ftz_f(value, shuffled) + warp_maxima += (value,) + if stage_info.loop_offset == 0: + if warp == 0 and lane < 8: + scratch[lane] = max_encode(cl.float32(NEG_MAX)) + cl.barrier_sync_block(number_of_threads=128, barrier_id=inst_idx + 1) + if lane < 8: + for pair in cl.static_iter(range(2)): + cl.atomic_rmw( + cl.AtomicOp.MAX, scratch.pointer((lane % 4) * 2 + pair), + max_encode(warp_maxima[pair]), + ) + else: + for pair in cl.static_iter(range(2)): + value = old_max[pair] + for k in cl.static_iter(range(4)): + value = cl.maximum(value, scores[2 * k + pair]) + shuffled, _ = cl.shuffle_sync(cl.ShuffleKind.XOR, value, 16) + value = cl.maximum(value, shuffled) + shuffled, _ = cl.shuffle_sync(cl.ShuffleKind.XOR, value, 8) + value = cl.maximum(value, shuffled) + if lane < 8: + cl.atomic_rmw( + cl.AtomicOp.MAX, scratch.pointer((lane % 4) * 2 + pair), max_encode(value) + ) + cl.barrier_sync_block(number_of_threads=128, barrier_id=inst_idx + 1) + for pair in cl.static_iter(range(2)): + maxima += (max_decode(scratch[(lane % 4) * 2 + pair]),) + return scores, cl.Vector(*maxima) + + +@dataclass(kw_only=True, eq=False) +class SmemPResource(ts.MemoryResource): + @ts.producer_work(outputs=1) + @staticmethod + def compute_p(stage_info, scores, new_max, offset, cfg): + warp, lane = warp_and_lane(stage_info) + probabilities = () + for k in cl.static_iter(range(4)): + score_pair = cl.Vector(scores[2 * k], scores[2 * k + 1]) + shifted = cl.sub(score_pair, new_max, rounding_mode=cl.RoundingMode.RN) + p_pair = cl.exp2( + shifted * pair_vector(128**-0.5 * 1.4426950408889634), flush_to_zero=True + ) + probabilities += ( + p_pair[0] if new_max[0] != NEG_MAX else cl.float32(0), + p_pair[1] if new_max[1] != NEG_MAX else cl.float32(0), + ) + local_sum = pair_vector(0) + for k in cl.static_iter((0, 2, 1, 3)): + local_sum = cl.add( + local_sum, + cl.Vector(probabilities[2 * k], probabilities[2 * k + 1]), + rounding_mode=cl.RoundingMode.RN, + ) + registers = () + for k in cl.static_iter(range(4)): + packed = cl.Vector( + cfg.q_dtype(probabilities[2 * k]), cfg.q_dtype(probabilities[2 * k + 1]) + ) + registers += (packed.reinterpret_as_scalar(cl.int32),) + smem = smem_array(stage_info, offset, cl.int32, 512) + row, matrix = lane % 8, lane // 8 + col = (warp % 2) * 4 + matrix + byte_offset = (warp // 2) * 1024 + row * 128 + ((col ^ row) * 16) + cl.store_matrix( + smem.pointer(byte_offset // 4), + cl.Vector(*registers), + shape=cl.MatrixStoreShape.M8N8, + transpose=True, + ) + cl.fence_proxy_bidirectional( + cl.FenceProxy.ASYNC, restriction=cl.FenceRestriction.shared_block() + ) + return local_sum + + @ts.consumer_work(outputs=1) + @staticmethod + def p_desc(stage_info, offset, cfg): + return smem_descriptor(smem_array(stage_info, offset, cfg.q_dtype, 1024).pointer(), 1024) + + +@dataclass(kw_only=True, eq=False) +class TmemSoftmaxStatsResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def store_stats(stage_info, values, maxima, column): + warp, _ = warp_and_lane(stage_info) + cl.tcgen05_store( + cl.Tcgen05LoadStoreShape.SHAPE_32X32B, + tmem_ptr(stage_info, column, warp * 32), + cl.Vector(values[0], values[1], maxima[0], maxima[1]), + ) + cl.tcgen05_wait_store() + + @ts.consumer_work(outputs=1) + @staticmethod + def load_stats(stage_info, column): + warp, _ = warp_and_lane(stage_info) + values = cl.tcgen05_load( + cl.Tcgen05LoadStoreShape.SHAPE_32X32B, + tmem_ptr(stage_info, column, warp * 32), + element_count=4, + dtype=cl.float32, + ) + cl.tcgen05_wait_load() + return values + + +@dataclass(kw_only=True, eq=False) +class TmemOResource(ts.MemoryResource): + @ts.producer_work + @staticmethod + def pv_mma(stage_info, v_desc, p_desc, section, cfg): + instr = cl.Tcgen05InstructionDescriptor( + d_type=cl.float32, + a_type=cfg.q_dtype, + b_type=cfg.q_dtype, + m=128, + n=8, + transpose_a=True, + ).encode() + accumulate = stage_info.loop_offset > 0 if section == 1 else num_pairs(stage_info, cfg) > 1 + if cl.elect_sync(): + for ki in cl.static_iter(range(8)): + cl.tcgen05_mma( + cl.Tcgen05MMAKind.F16, + tmem_ptr(stage_info, 80 + stage_info.stage_idx * 8), + v_desc + ki * 128, + p_desc + ki * 2 + (ki // 4) * 56, + instr, + accumulate=accumulate if ki == 0 else True, + cta_group=cl.CTAGroup.CTA_1, + ) + + @ts.consumer_work + @staticmethod + def correction(stage_info, stats): + unchanged = (stats[0] == stats[2]) & (stats[1] == stats[3]) + if not cl.vote_all_sync(unchanged): + scales = cl.Vector( + cl.float32(1) if stats[0] == stats[2] else rescale(stats[0], stats[2]), + cl.float32(1) if stats[1] == stats[3] else rescale(stats[1], stats[3]), + ) + column = 80 + stage_info.stage_idx * 8 + values = load_o(stage_info, column) + corrected = values * cl.Vector(*tuple(scales[i % 2] for i in cl.static_iter(range(8)))) + store_o(stage_info, column, corrected) + cl.tcgen05_wait_store() + + @ts.consumer_work + @staticmethod + def tail_epilogue(stage_info, stats0, stats1, scratch_offset, output_offset, cfg): + warp, lane = warp_and_lane(stage_info) + group, kv_head, batch = work_coords(stage_info) + inputs = stage_info.context.tasks_inputs + scratch = smem_array(stage_info, scratch_offset, cl.float32, 32) + exp0, exp1 = (), () + for pair in cl.static_iter(range(2)): + maximum = cl.maximum(stats0[pair + 2], stats1[pair + 2]) + if cfg.use_variable_seqlens_q: + maximum = cl._nvvm.fmax_ftz_f(stats0[pair + 2], stats1[pair + 2]) + e0 = rescale(stats0[pair + 2], maximum) + e1 = rescale(stats1[pair + 2], maximum) + exp0 += (e0,) + exp1 += (e1,) + total = stats0[pair] * e0 + stats1[pair] * e1 + for mask in cl.static_iter((16, 8, 4)): + shuffled, _ = cl.shuffle_sync(cl.ShuffleKind.XOR, total, mask) + total += shuffled + if lane < 4: + scratch[warp * 8 + lane * 2 + pair] = total + cl.barrier_sync_block(number_of_threads=128, barrier_id=3) + denom = () + for pair in cl.static_iter(range(2)): + total = cl.float32(0) + for w in cl.static_iter(range(4)): + total += scratch[w * 8 + (lane % 4) * 2 + pair] + denom += (total,) + o0, o1 = load_o(stage_info, 80), load_o(stage_info, 88) + inverse = cl.Vector(safe_norm_rcp(denom[0]), safe_norm_rcp(denom[1])) + scale0 = cl.Vector(*exp0) * inverse + scale1 = cl.Vector(*exp1) * inverse + registers = () + for k in cl.static_iter(range(4)): + pair0 = cl.Vector(o0[2 * k], o0[2 * k + 1]) + pair1 = cl.Vector(o1[2 * k], o1[2 * k + 1]) + final_pair = cl.fma(pair0, scale0, pair1 * scale1) + registers += (final_pair.astype(cfg.out_dtype).reinterpret_as_scalar(cl.int32),) + smem = smem_array(stage_info, output_offset, cl.int32, 512) + row = lane % 8 + col = (warp % 2) * 4 + lane // 8 + byte_offset = (warp // 2) * 1024 + row * 128 + ((col ^ row) * 16) + cl.store_matrix( + smem.pointer(byte_offset // 4), + cl.Vector(*registers), + shape=cl.MatrixStoreShape.M8N8, + transpose=True, + ) + cl.fence_proxy_bidirectional( + cl.FenceProxy.ASYNC, restriction=cl.FenceRestriction.shared_block() + ) + cl.barrier_sync_block(number_of_threads=128, barrier_id=3) + base_offset = (warp * 32 + lane) * 16 + smem_row = base_offset // 128 + load_offset = base_offset ^ ((smem_row % 8) * 16) + head = group * 8 + smem_row % 8 + q_valid = head < inputs.head_ratio + if cfg.use_variable_seqlens_q: + token, head = q_row_token_and_local_head(group, smem_row % 8, cfg) + batch = inputs.q_token_offset + token + q_valid = (token < inputs.seq_len_q) & (head < cfg.heads_q_per_kv) + column = (smem_row // 8) * 64 + (base_offset % 128) // 2 + if q_valid: + packed = smem.pointer(load_offset // 4).load(count=4, alignment=16) + dst = inputs.o.pointer((batch, kv_head * inputs.head_ratio + head, column)) + cl.bitcast(dst, cl.pointer_dtype(cl.int32, cl.MemorySpace.GLOBAL)).store( + packed, alignment=16 + ) diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_run.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_run.py new file mode 100644 index 00000000..7a8d3893 --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_run.py @@ -0,0 +1,135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Paged HND decode launch for SQ1 or packed Q and a PyTorch reference.""" + +import itertools + +import cuda.lang as cl +import torch + +from .fmha_decode_config import FmhaDecodeConfig, validate_config +from .fmha_decode_kernel import make_fmha_decode_launcher + + +def prepare_tensors( + lengths, heads_q=8, heads_kv=1, cfg=FmhaDecodeConfig(), seed=1111, *, q_lengths=None +): + validate_config(cfg) + if not lengths or min(lengths) < 1: + raise ValueError("decode requires at least one K/V token per request") + if cfg.use_variable_seqlens_q: + if q_lengths is None or len(q_lengths) != len(lengths): + raise ValueError("packed Q requires one query length per request") + if min(q_lengths) < 1 or max(q_lengths) > cfg.max_seq_len_q: + raise ValueError("query lengths must be in [1,max_seq_len_q]") + if cfg.mask_type == "causal" and any(q > k for q, k in zip(q_lengths, lengths)): + raise ValueError("causal attention requires q_len <= kv_len for each request") + elif q_lengths is not None: + raise ValueError("q_lengths requires use_variable_seqlens_q=True") + torch.manual_seed(seed) + page = cfg.num_tokens_per_page + counts = [(n + page - 1) // page for n in lengths] + shape = (sum(counts), heads_kv, page, 128) + dtype = torch.bfloat16 if cfg.q_dtype == cl.bfloat16 else torch.float16 + total_q = sum(q_lengths) if cfg.use_variable_seqlens_q else len(lengths) + tensors = dict( + q=(torch.randn(total_q, heads_q, 128, device="cuda") * 0.2).to(dtype), + k=(torch.randn(shape, device="cuda") * 0.2).to(dtype), + v=(torch.randn(shape, device="cuda") * 0.2).to(dtype), + o=torch.empty(total_q, heads_q, 128, device="cuda", dtype=dtype), + seq_lens=torch.tensor(lengths, device="cuda", dtype=torch.int32), + paged_kv_indptr=torch.tensor( + [0] + list(itertools.accumulate(counts)), device="cuda", dtype=torch.int32 + ), + paged_kv_indices=torch.arange(sum(counts) - 1, -1, -1, device="cuda", dtype=torch.int32), + ) + if cfg.use_variable_seqlens_q: + tensors["qo_indptr"] = torch.tensor( + [0] + list(itertools.accumulate(q_lengths)), device="cuda", dtype=torch.int32 + ) + return tensors + + +def run(tensors, cfg=FmhaDecodeConfig(), stream=None): + validate_config(cfg) + q, k, v, o = (tensors[name] for name in ("q", "k", "v", "o")) + dtype = torch.bfloat16 if cfg.q_dtype == cl.bfloat16 else torch.float16 + if any( + t.dtype != dtype or not t.is_cuda or not t.is_contiguous() for t in (q, k, v, o) + ): + raise ValueError(f"Q/K/V/O must be contiguous CUDA tensors with config dtype {dtype}") + if q.ndim != 3 or k.ndim != 4 or q.shape[-1] != 128 or k.shape[-1] != 128: + raise ValueError("expected Q[B or total_q,Hq,128] and K/V[pages,Hkv,page,128]") + if min(q.shape) < 1 or min(k.shape) < 1: + raise ValueError("Q and K/V tensor extents must be positive") + if v.shape != k.shape or o.shape != q.shape or k.shape[2] != cfg.num_tokens_per_page: + raise ValueError("inconsistent Q/K/V/O geometry") + if q.shape[1] % k.shape[1] or not 1 <= q.shape[1] // k.shape[1] <= 32: + raise ValueError("Hq/Hkv must be an integer in [1,32]") + metadata = tuple(tensors[name] for name in ("seq_lens", "paged_kv_indptr", "paged_kv_indices")) + if cfg.use_variable_seqlens_q: + if "qo_indptr" not in tensors: + raise ValueError("packed Q requires qo_indptr") + metadata += (tensors["qo_indptr"],) + elif "qo_indptr" in tensors: + raise ValueError("qo_indptr requires use_variable_seqlens_q=True") + if any(t.device != q.device for t in (k, v, o, *metadata)): + raise ValueError("all tensors must reside on the same CUDA device") + if any(t.dtype != torch.int32 or t.ndim != 1 or not t.is_contiguous() for t in metadata): + raise ValueError("metadata must be contiguous one-dimensional int32 tensors") + batch = metadata[0].numel() + if batch < 1 or metadata[1].numel() != batch + 1: + raise ValueError("sequence lengths and page indptr must match the request batch") + if cfg.use_variable_seqlens_q: + if metadata[3].numel() != batch + 1: + raise ValueError("qo_indptr must have B+1 entries") + if not batch <= q.shape[0] <= batch * cfg.max_seq_len_q: + raise ValueError("packed query extent must be between B and B*max_seq_len_q") + elif batch != q.shape[0]: + raise ValueError("sequence lengths must match the SQ1 query batch") + if batch * cfg.max_seq_len_q * q.shape[1] > 2147483647: + raise ValueError("planned query/head extent must fit in int32") + ratio = q.shape[1] // k.shape[1] + launcher = make_fmha_decode_launcher(ratio, k.shape[1], cfg, batch=batch) + args = tuple( + tensors[name] + for name in ("q", "k", "v", "o", "seq_lens", "paged_kv_indptr", "paged_kv_indices") + ) + if cfg.use_variable_seqlens_q: + args += (tensors["qo_indptr"],) + launcher(torch.cuda.current_stream() if stream is None else stream, *args) + return o + + +def torch_reference(tensors, cfg=FmhaDecodeConfig()): + q, k, v = (tensors[name] for name in ("q", "k", "v")) + result = torch.empty_like(q) + indptr = tensors["paged_kv_indptr"].tolist() + ratio = q.shape[1] // k.shape[1] + q_offsets = ( + tensors["qo_indptr"].tolist() if cfg.use_variable_seqlens_q + else list(range(q.shape[0] + 1)) + ) + for batch, length in enumerate(tensors["seq_lens"].tolist()): + ids = tensors["paged_kv_indices"][indptr[batch]:indptr[batch + 1]].long() + q_begin, q_end = q_offsets[batch:batch + 2] + q_length = q_end - q_begin + keys = torch.arange(length, device=q.device) + rows = torch.arange(q_length, device=q.device) + upper = length - q_length + rows + 1 + visible = torch.ones(q_length, length, device=q.device, dtype=torch.bool) + if cfg.mask_type == "causal": + visible &= keys[None, :] < upper[:, None] + if cfg.window_left >= 0: + visible &= keys[None, :] >= upper[:, None] - cfg.window_left - 1 + for head in range(k.shape[1]): + kk = k[ids, head].reshape(-1, 128)[:length].float() + vv = v[ids, head].reshape(-1, 128)[:length].float() + qq = q[q_begin:q_end, head * ratio:(head + 1) * ratio].float() + scores = qq @ kk.T / 128**0.5 + scores.masked_fill_(~visible[:, None, :], float("-inf")) + result[q_begin:q_end, head * ratio:(head + 1) * ratio] = ( + torch.softmax(scores, -1) @ vv + ).to(q.dtype) + return result diff --git a/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_tasks.py b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_tasks.py new file mode 100644 index 00000000..52fe84dd --- /dev/null +++ b/experimental/task_scheduling/tutorial/08_fmha_decode_sm100a/fmha_decode_tasks.py @@ -0,0 +1,478 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-License-Identifier: Apache-2.0 + +"""Two alternating K/V instances with peeled QK head and PV tail.""" + +from contextlib import contextmanager +from dataclasses import dataclass, field, replace + +import cuda.lang as cl +import task_scheduling as ts +from task_scheduling.task import DeviceWorkTileLoop + +from .fmha_decode_config import FmhaDecodeConfig +from .fmha_decode_resources.helpers_common import kv_tile_bounds, q_group_token_base + + +@dataclass(frozen=True) +class _ResolvedPackedDecodeWorkQueue: + cfg: FmhaDecodeConfig + qo_indptr: object + + +@dataclass(frozen=True) +class _DevicePackedDecodeWorkQueue: + cfg: FmhaDecodeConfig + + def bind_inputs(self, inputs): + return _ResolvedPackedDecodeWorkQueue(self.cfg, inputs.qo_indptr) + + +@dataclass(kw_only=True, eq=False) +class PackedDecodeWorkQueue(ts.WorkQueue): + """Skip inactive packed-Q tiles without skipping scheduler handoffs.""" + + cfg: FmhaDecodeConfig = field(init=False, default=None) + + def __init__(self, cfg, **kwargs): + super().__init__(**kwargs) + self.cfg = cfg + + def _freeze_skip_context(self): + return _DevicePackedDecodeWorkQueue(self.cfg) + + def skip_work_tile_if(self, work_tile): + group, _, batch = work_tile.tile_idx + length = self.qo_indptr[batch + 1] - self.qo_indptr[batch] + return q_group_token_base(group, self.cfg) >= length + + +@dataclass(frozen=True) +class _RefreshDecodeInputs: + """Bind request metadata once per logical tile, before its pipeline work.""" + + packed_q: bool + + def __call__(self, context): + inputs = context.tasks_inputs + batch = context.work_tile.tile_idx[2] + begin = inputs.paged_kv_indptr[batch] + q_begin, q_length = inputs.q_token_offset, inputs.seq_len_q + if self.packed_q: + q_begin = inputs.qo_indptr[batch] + q_length = inputs.qo_indptr[batch + 1] - q_begin + return replace( + context, + tasks_inputs=replace( + inputs, + seq_len=inputs.seq_lens[batch], + page_begin=begin, + page_count=inputs.paged_kv_indptr[batch + 1] - begin, + q_token_offset=q_begin, + seq_len_q=q_length, + ), + ) + + +@dataclass(kw_only=True, eq=False) +class ScheduleTokenThrottleResource(ts.MemoryResource): + """Prevent the scheduler from recycling a token before Load owns it.""" + + @ts.producer_work + @staticmethod + def publish_schedule_token(stage_info): + pass + + @ts.consumer_work + @staticmethod + def consume_schedule_token(stage_info): + pass + + +def _work_queue_tail(work_queue): + work_queue.wait() + work_queue.get_and_advance_work_tile() + work_queue.release() + + +def _schedule_token_throttle_head(throttle): + if throttle is not None: + throttle.acquire() + throttle.publish_schedule_token() + throttle.commit() + + +@contextmanager +def _work_tile_schedule(work_queue, cfg, throttle=None): + if work_queue is None: + yield + elif cfg.use_variable_seqlens_q: + with ts.work_tile_loop( + work_queue, skip_if=PackedDecodeWorkQueue.skip_work_tile_if + ) as work_tiles: + _schedule_token_throttle_head(throttle) + with work_tiles.skippable(): + yield + _work_queue_tail(work_queue) + else: + with ts.work_tile_loop(work_queue): + _schedule_token_throttle_head(throttle) + yield + _work_queue_tail(work_queue) + + +@dataclass(frozen=True) +class _ResolvedDecodeDomain: + seq_len: object + seq_len_q: object + offset: int + cfg: FmhaDecodeConfig + + +@dataclass(frozen=True) +class _DeviceDecodeDomain: + offset: int + cfg: FmhaDecodeConfig + + def bind_inputs(self, inputs): + return _ResolvedDecodeDomain(inputs.seq_len, inputs.seq_len_q, self.offset, self.cfg) + + +class DecodeDomainTask(ts.Task): + def __init__(self, *args, cfg, offset=0, **kwargs): + super().__init__(*args, **kwargs) + self.offset = offset + self.cfg = cfg + + def _freeze_domain_task(self): + return _DeviceDecodeDomain(self.offset, self.cfg) + + def to_device(self, *args, **kwargs): + device_task = super().to_device(*args, **kwargs) + # Keep request metadata in SSA values across all resource/domain steps. + # Pipeline phases remain live when a persistent CTA takes another tile. + return replace( + device_task, + body=tuple( + replace( + node, body=(_RefreshDecodeInputs(self.cfg.use_variable_seqlens_q),) + node.body + ) + if isinstance(node, DeviceWorkTileLoop) + else node + for node in device_task.body + ), + ) + + def get_domain(self, tile_coord): + start, end = kv_tile_bounds(self.seq_len, self.seq_len_q, tile_coord[0], self.cfg) + if self.cfg.use_variable_seqlens_q: + return ((end - start + 1) >> 1) - self.offset + return (end - start + 1) // 2 - self.offset + + def get_last_iteration(self, tile_coord): + return cl.maximum(DecodeDomainTask.get_domain(self, tile_coord) - 1, 0) + + +def create_load_task(sq, skv, pages, offsets, cfg, work_queue=None, throttle=None): + def schedule_body(sq, skv, pages, wq, throttle): + with _work_tile_schedule(wq, cfg, throttle): + + def load(inst, section, is_v): + pages.wait() + ids = pages.page_ids(offsets["pages"], inst, section, is_v, cfg) + skv.acquire() + skv.tma_load(ids, offsets["kv"], is_v, cfg) + skv.commit() + pages.release() + + sq.acquire() + sq.tma_load(offsets["q"], cfg) + sq.commit() + load(0, 0, False) + load(1, 0, False) + + def body(): + load(0, 1, False) + load(0, 1, True) + load(1, 1, False) + load(1, 1, True) + + ts.domain_loop(0, DecodeDomainTask.get_domain, 1, body) + load(0, 2, True) + load(1, 2, True) + + @ts.schedule + def direct_schedule(stage_info, sq, skv, pages): + schedule_body(sq, skv, pages, None, None) + + @ts.schedule + def persistent_schedule(stage_info, sq, skv, pages, wq, throttle): + schedule_body(sq, skv, pages, wq, throttle) + + captured = ( + persistent_schedule(sq, skv, pages, work_queue, throttle) + if work_queue is not None + else direct_schedule(sq, skv, pages) + ) + + return DecodeDomainTask( + name="LoadTask", + warp_idx=cfg.clc_load_warp_idx if cfg.use_persistent_scheduler else cfg.load_warp_idx, + num_warps=1, + schedule=captured, + offset=1, + cfg=cfg, + ) + + +def create_page_offsets_task(pages, offset, cfg, work_queue=None): + def schedule_body(pages, wq): + with _work_tile_schedule(wq, cfg): + def load(inst, section, is_v): + pages.acquire() + pages.load_page_offsets(offset, inst, section, is_v, cfg) + pages.commit() + + load(0, 0, False) + load(1, 0, False) + + def body(): + load(0, 1, False) + load(0, 1, True) + load(1, 1, False) + load(1, 1, True) + + ts.domain_loop(0, DecodeDomainTask.get_domain, 1, body) + load(0, 2, True) + load(1, 2, True) + + @ts.schedule + def direct_schedule(stage_info, pages): + schedule_body(pages, None) + + @ts.schedule + def persistent_schedule(stage_info, pages, wq): + schedule_body(pages, wq) + + captured = ( + persistent_schedule(pages, work_queue) + if work_queue is not None + else direct_schedule(pages) + ) + + return DecodeDomainTask( + name="PageOffsetsTask", + warp_idx=14, + num_warps=1, + schedule=captured, + offset=1, + cfg=cfg, + ) + + +def create_mma_task(sq, skv, s0, s1, p0, p1, o, offsets, cfg, work_queue=None): + def schedule_body(sq, skv, s0, s1, p0, p1, o, wq): + with _work_tile_schedule(wq, cfg): + sq.wait() + q_desc = sq.q_desc(offsets["q"], cfg) + + def qk(s, column): + s.acquire() + skv.wait() + k_desc = skv.kv_desc(offsets["kv"], cfg) + s.qk_mma(q_desc, k_desc, column, cfg) + skv.release() + s.commit() + + def pv(p, p_offset, section): + p.wait() + p_desc = p.p_desc(p_offset, cfg) + o.acquire() + skv.wait() + v_desc = skv.kv_desc(offsets["kv"], cfg) + o.pv_mma(v_desc, p_desc, section, cfg) + skv.release() + o.commit() + p.release() + + qk(s0, 0) + qk(s1, 8) + + def body(): + qk(s0, 0) + pv(p0, offsets["p0"], 1) + qk(s1, 8) + pv(p1, offsets["p1"], 1) + + ts.domain_loop(0, DecodeDomainTask.get_domain, 1, body) + pv(p0, offsets["p0"], 2) + pv(p1, offsets["p1"], 2) + sq.release() + + @ts.schedule + def direct_schedule(stage_info, sq, skv, s0, s1, p0, p1, o): + schedule_body(sq, skv, s0, s1, p0, p1, o, None) + + @ts.schedule + def persistent_schedule(stage_info, sq, skv, s0, s1, p0, p1, o, wq): + schedule_body(sq, skv, s0, s1, p0, p1, o, wq) + + captured = ( + persistent_schedule(sq, skv, s0, s1, p0, p1, o, work_queue) + if work_queue is not None + else direct_schedule(sq, skv, s0, s1, p0, p1, o) + ) + + return DecodeDomainTask( + name="MmaTask", + warp_idx=12, + num_warps=1, + schedule=captured, + offset=1, + cfg=cfg, + ) + + +def create_softmax_task(s, p, stats, inst, offsets, cfg, work_queue=None): + def schedule_body(s, p, stats, wq): + with _work_tile_schedule(wq, cfg): + maximum, total = s.init_state() + + def body(maximum, total): + s.wait() + scores, new_max = s.compute_softmax_loop( + maximum, inst * 8, offsets["max" + str(inst)], inst, cfg + ) + s.release() + stats.acquire() + stats.store_stats(maximum, new_max, 16 + inst * 32) + stats.commit() + p.acquire() + local_sum = p.compute_p(scores, new_max, offsets["p" + str(inst)], cfg) + p.commit() + total = s.reduce_sums(maximum, new_max, total, local_sum) + return new_max, total + + # Peel the final iteration to keep its stats handoff out of the steady-state loop. + maximum, total = ts.domain_loop( + 0, DecodeDomainTask.get_last_iteration, 1, body, maximum, total, unroll=1 + ) + + def last_body(maximum, total): + new_max, total = body(maximum, total) + with ts.last_iter(): + stats.acquire() + stats.store_stats(total, new_max, 16 + inst * 32) + stats.commit() + return new_max, total + + ts.domain_loop( + DecodeDomainTask.get_last_iteration, + DecodeDomainTask.get_domain, + 1, + last_body, + maximum, + total, + unroll=1, + ) + + @ts.schedule + def direct_schedule(stage_info, s, p, stats): + schedule_body(s, p, stats, None) + + @ts.schedule + def persistent_schedule(stage_info, s, p, stats, wq): + schedule_body(s, p, stats, wq) + + captured = ( + persistent_schedule(s, p, stats, work_queue) + if work_queue is not None + else direct_schedule(s, p, stats) + ) + + return DecodeDomainTask( + name="Softmax" + str(inst) + "Task", + warp_idx=inst * 4, + num_warps=4, + schedule=captured, + cfg=cfg, + ) + + +def create_correction_task(stats0, stats1, o, offsets, cfg, work_queue=None): + def schedule_body(stats0, stats1, o, wq): + with _work_tile_schedule(wq, cfg): + stats0.wait() + stats0.release() + stats1.wait() + stats1.release() + + def correct(stats, column): + stats.wait() + values = stats.load_stats(column) + stats.release() + o.wait() + o.correction(values) + o.release() + + def body(): + correct(stats0, 16) + correct(stats1, 48) + + ts.domain_loop(0, DecodeDomainTask.get_domain, 1, body) + stats0.wait() + values0 = stats0.load_stats(16) + stats0.release() + o.wait() + stats1.wait() + values1 = stats1.load_stats(48) + stats1.release() + o.wait() + o.tail_epilogue(values0, values1, offsets["sum"], offsets["output"], cfg) + o.release() + o.release() + + @ts.schedule + def direct_schedule(stage_info, stats0, stats1, o): + schedule_body(stats0, stats1, o, None) + + @ts.schedule + def persistent_schedule(stage_info, stats0, stats1, o, wq): + schedule_body(stats0, stats1, o, wq) + + captured = ( + persistent_schedule(stats0, stats1, o, work_queue) + if work_queue is not None + else direct_schedule(stats0, stats1, o) + ) + + return DecodeDomainTask( + name="CorrectionTask", + warp_idx=8, + num_warps=4, + schedule=captured, + offset=1, + cfg=cfg, + ) + + +def create_scheduler_task(work_queue, throttle, cfg): + @ts.schedule + def schedule(stage_info, wq, throttle): + with ts.work_tile_loop(wq): + with ts.domain_loop(0): + pass + throttle.wait() + throttle.consume_schedule_token() + throttle.release() + wq.acquire() + wq.fetch_work_tile() + wq.commit() + _work_queue_tail(wq) + + return ts.Task( + name="SchedulerTask", + warp_idx=cfg.scheduler_warp_idx, + num_warps=1, + schedule=schedule(work_queue, throttle), + ) From 49671db48d8cfeed8761e0d0c55a019902f94a30 Mon Sep 17 00:00:00 2001 From: Asher Mancinelli Date: Wed, 7 Oct 2026 06:52:03 -0700 Subject: [PATCH 2/7] Move DType definition to C++ Co-authored-by: Greg Bonik Signed-off-by: Asher Mancinelli --- cext/CMakeLists.txt | 1 + cext/dtype.cpp | 687 ++++++++++++++++++ cext/dtype.h | 289 ++++++++ cext/module.cpp | 4 + cext/py.h | 233 +++--- .../cuda-lang/src/cuda/lang/_datatype.py | 29 +- experimental/cuda-lang/test/test_dtype.py | 18 + src/cuda/tile/_cext.pyi | 69 +- src/cuda/tile/_datatype.py | 354 ++------- test/test_ir_types.py | 17 +- 10 files changed, 1290 insertions(+), 411 deletions(-) create mode 100644 cext/dtype.cpp create mode 100644 cext/dtype.h diff --git a/cext/CMakeLists.txt b/cext/CMakeLists.txt index 7cebcb49..5380f53c 100644 --- a/cext/CMakeLists.txt +++ b/cext/CMakeLists.txt @@ -60,6 +60,7 @@ add_library(_cext_static STATIC coroutine_util.cpp cuda_loader.cpp cuda_helper.cpp + dtype.cpp ipc_util.cpp memory.cpp py.cpp diff --git a/cext/dtype.cpp b/cext/dtype.cpp new file mode 100644 index 00000000..2e355357 --- /dev/null +++ b/cext/dtype.cpp @@ -0,0 +1,687 @@ +// SPDX-FileCopyrightText: Copyright (c) <2026> NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// +// SPDX-License-Identifier: Apache-2.0 + +#include "dtype.h" +#include "vec.h" +#include "py.h" + +static PyObject* g_module; +static PyObject* g_pointee_dtype_pyunicode; +static PyObject* g_memory_space_pyunicode; + + +int append_to_string_builder(Integer integer, StringBuilder* sb) { + if (integer.is_signed) { + sb->append(static_cast(integer.bits)); + } else { + sb->append(static_cast(integer.bits)); + } + return 0; +} + + +// ----- Memory space implementation ----- + +const char* memory_space_str(MemorySpace space) { + switch (space) { + #define MEMORY_SPACE_STR_CASE(name, _id, _ptrwidth) \ + case MemorySpace::name: return #name; + FOREACH_MEMORY_SPACE(MEMORY_SPACE_STR_CASE) + #undef MEMORY_SPACE_STR_CASE + } + CHECK_UNREACHABLE; +} + +static constexpr MemorySpace all_memory_spaces[] = { + #define MEMORY_SPACE_ENUMERATE(name, _id, _ptrwidth) \ + MemorySpace::name, + FOREACH_MEMORY_SPACE(MEMORY_SPACE_ENUMERATE) + #undef MEMORY_SPACE_ENUMERATE +}; + +static PyObject* const* get_memory_space_pyobjects(GlobalLock&) { + static bool is_cached; + static PyObject* cache[kMemorySpaceMax + 1]; + if (!is_cached) { + PyPtr mod = steal(PyImport_ImportModule("cuda.tile._memory_model")); + if (!mod) return nullptr; + PyPtr cls = getattr(mod, "MemorySpace"); + if (!cls) return nullptr; + + for (MemorySpace space : all_memory_spaces) { + uint8_t i = static_cast(space); + if (cache[i]) continue; + cache[i] = getattr(cls, memory_space_str(space)).release(); + if (!cache[i]) return nullptr; + } + is_cached = true; + } + return cache; +} + +static PyObject* memory_space_to_pyobject(MemorySpace space, GlobalLock& lock) { + PyObject* const* objects = get_memory_space_pyobjects(lock); + if (!objects) return nullptr; + return objects[static_cast(space)]; +} + +static std::optional memory_space_from_pyobject(PyObject* space_obj, + GlobalLock& lock) { + PyObject* const* objects = get_memory_space_pyobjects(lock); + for (MemorySpace space : all_memory_spaces) { + if (objects[static_cast(space)] == space_obj) + return space; + } + return std::nullopt; +} + + +// ----- DType implementation ----- + +namespace _dtype_detail { + const uint32_t basic_dtype_bitwidth[basic_last + 1] = { + 0, + #define BASIC_DTYPE_BITWIDTH(_name, bitwidth, _signed, _doc) \ + bitwidth, + FOREACH_BASIC_DTYPE(BASIC_DTYPE_BITWIDTH) + #undef BASIC_DTYPE_BITWIDTH + }; + const bool basic_dtype_is_signed[basic_last + 1] = { + false, + #define BASIC_DTYPE_IS_SIGNED(_name, _bitwidth, signed, _doc) \ + signed, + FOREACH_BASIC_DTYPE(BASIC_DTYPE_IS_SIGNED) + #undef BASIC_DTYPE_IS_SIGNED + }; +} // namespace _dtype_detail + +enum class DerivedDTypeKind : uint8_t { + Pointer = 1, + ForeignPointer, +}; + +struct DerivedDType { + DerivedDTypeKind kind; + MemorySpace memory_space; // when kind is `Pointer` + DType pointee_dtype; // when kind is `Pointer` or `ForeignPointer` +}; + +struct DTypeRegistry { + Vec pointer_dtypes[kMemorySpaceMax + 1]; + Vec foreign_pointer_dtypes; + Vec derived_dtypes; + Vec names; + Vec pyobjects; + + static DTypeRegistry& get(GlobalLock& lock) { + DTypeRegistry*& ptr = instance_.get(lock); + if (!ptr) ptr = new DTypeRegistry; + return *ptr; + } + + static ProtectedByGlobalLock instance_; +private: + DTypeRegistry() = default; + ~DTypeRegistry(); // DTypeRegistry is an immortal singleton +}; + +ProtectedByGlobalLock DTypeRegistry::instance_; + + +Integer integer_dtype_min(DType dtype) { + if (is_signed_integer_dtype(dtype)) { + uint32_t width = _dtype_detail::basic_dtype_bitwidth[dtype.dtype_id]; + CHECK(width <= 64); + uint64_t bits = static_cast(-1) << (width - 1); + return Integer::from_i64(static_cast(bits)); + } else { + CHECK(is_unsigned_integer_dtype(dtype)); + return Integer::from_u64(0); + } +} + +Integer integer_dtype_max(DType dtype) { + uint32_t width = _dtype_detail::basic_dtype_bitwidth[dtype.dtype_id]; + CHECK(width <= 64); + if (is_signed_integer_dtype(dtype)) { + uint64_t bits = (static_cast(1) << (width - 1)) - 1; + return Integer::from_i64(static_cast(bits)); + } else { + CHECK(is_unsigned_integer_dtype(dtype)); + uint64_t bits = (static_cast(-1)) >> (64 - width); + return Integer::from_u64(bits); + } +} + +DType pointer_dtype(DType pointee_dtype, MemorySpace memory_space, GlobalLock& lock) { + uint32_t space_idx = static_cast(memory_space); + CHECK(space_idx <= kMemorySpaceMax); + DTypeRegistry& reg = DTypeRegistry::get(lock); + Vec& pointer_dtypes = reg.pointer_dtypes[space_idx]; + + if (pointee_dtype.dtype_id >= pointer_dtypes.size()) + pointer_dtypes.resize(pointee_dtype.dtype_id + 1); + + DType& ret = pointer_dtypes[pointee_dtype.dtype_id]; + if (!ret) { + uint64_t dtype_id = static_cast(reg.derived_dtypes.size()) + kFirstDerivedDTypeId; + CHECK(dtype_id <= UINT32_MAX); + reg.derived_dtypes.push_back(DerivedDType{ + DerivedDTypeKind::Pointer, memory_space, pointee_dtype}); + ret.dtype_id = static_cast(dtype_id); + } + return ret; +} + +DType foreign_pointer_dtype(DType pointee_dtype, GlobalLock& lock) { + DTypeRegistry& reg = DTypeRegistry::get(lock); + if (pointee_dtype.dtype_id >= reg.foreign_pointer_dtypes.size()) + reg.foreign_pointer_dtypes.resize(pointee_dtype.dtype_id + 1); + + DType& ret = reg.foreign_pointer_dtypes[pointee_dtype.dtype_id]; + if (!ret) { + uint64_t dtype_id = static_cast(reg.derived_dtypes.size()) + kFirstDerivedDTypeId; + CHECK(dtype_id <= UINT32_MAX); + reg.derived_dtypes.push_back(DerivedDType{ + DerivedDTypeKind::ForeignPointer, MemorySpace::GENERIC, pointee_dtype}); + ret.dtype_id = static_cast(dtype_id); + } + return ret; +} + +static const DerivedDType& get_derived_dtype(DType dtype, GlobalLock& lock) { + DTypeRegistry& reg = DTypeRegistry::get(lock); + uint32_t derived_idx = dtype.dtype_id - kFirstDerivedDTypeId; + CHECK(derived_idx < reg.derived_dtypes.size()); + return reg.derived_dtypes[derived_idx]; +} + +bool _dtype_detail::is_derived_dtype_pointer(DType dtype, GlobalLock& lock) { + return get_derived_dtype(dtype, lock).kind == DerivedDTypeKind::Pointer; +} + +bool _dtype_detail::is_derived_dtype_foreign_pointer(DType dtype, GlobalLock& lock) { + return get_derived_dtype(dtype, lock).kind == DerivedDTypeKind::ForeignPointer; +} + +DType pointer_dtype_pointee(DType pointer_dtype, GlobalLock& lock) { + const DerivedDType& derived = get_derived_dtype(pointer_dtype, lock); + CHECK(derived.kind == DerivedDTypeKind::Pointer); + return derived.pointee_dtype; +} + +MemorySpace pointer_dtype_memory_space(DType pointer_dtype, GlobalLock& lock) { + const DerivedDType& derived = get_derived_dtype(pointer_dtype, lock); + CHECK(derived.kind == DerivedDTypeKind::Pointer); + return derived.memory_space; +} + +DType foreign_pointer_dtype_pointee(DType foreign_pointer_dtype, GlobalLock& lock) { + const DerivedDType& derived = get_derived_dtype(foreign_pointer_dtype, lock); + CHECK(derived.kind == DerivedDTypeKind::ForeignPointer); + return derived.pointee_dtype; +} + +uint32_t _dtype_detail::derived_dtype_bitwidth(DType dtype, GlobalLock& lock) { + const DerivedDType& derived = get_derived_dtype(dtype, lock); + switch (derived.kind) { + case DerivedDTypeKind::Pointer: + return memory_space_pointer_bitwidth(derived.memory_space); + case DerivedDTypeKind::ForeignPointer: + return 64; + } + CHECK_UNREACHABLE; +} + +static const char* basic_dtype_name(_dtype_detail::BasicDTypeEnum basic) { + switch (basic) { + case _dtype_detail::BasicDTypeEnum::Invalid: return ""; + + #define BASIC_DTYPE_NAME(name, _width, _signed, _doc) \ + case _dtype_detail::BasicDTypeEnum::name: return #name; + FOREACH_BASIC_DTYPE(BASIC_DTYPE_NAME); + #undef BASIC_DTYPE_NAME + } + CHECK_UNREACHABLE; +} + +static const char* basic_dtype_doc(_dtype_detail::BasicDTypeEnum basic) { + switch (basic) { + case _dtype_detail::BasicDTypeEnum::Invalid: return nullptr; + + #define BASIC_DTYPE_NAME(name, _width, _signed, doc) \ + case _dtype_detail::BasicDTypeEnum::name: return doc; + FOREACH_BASIC_DTYPE(BASIC_DTYPE_NAME); + #undef BASIC_DTYPE_NAME + } + CHECK_UNREACHABLE; +} + +static PyPtr make_dtype_name(DType dtype, GlobalLock& lock) { + if (dtype.dtype_id < kFirstDerivedDTypeId) { + return steal(PyUnicode_FromString(basic_dtype_name( + static_cast<_dtype_detail::BasicDTypeEnum>(dtype.dtype_id)))); + } + + const DerivedDType& derived = get_derived_dtype(dtype, lock); + switch (derived.kind) { + case DerivedDTypeKind::Pointer: + { + StringBuilder builder; + DType pointee = pointer_dtype_pointee(dtype, lock); + int num_params = 0; + if (pointee) { + CHECK(pointee.dtype_id < dtype.dtype_id); + builder.append("pointer["); + builder.append(pointee); + num_params = 1; + } else { + builder.append("opaque_pointer"); + } + + MemorySpace space = pointer_dtype_memory_space(dtype, lock); + if (space != MemorySpace::GENERIC) { + builder.append(++num_params == 1 ? "[" : ", "); + builder.append("MemorySpace."); + builder.append(memory_space_str(space)); + } + + if (num_params) builder.append("]"); + return builder.build(); + } + case DerivedDTypeKind::ForeignPointer: + { + DType pointee = foreign_pointer_dtype_pointee(dtype, lock); + return to_pyunicode("foreign_pointer[", pointee, "]"); + } + } + CHECK_UNREACHABLE; +} + +PyObject* dtype_name(DType dtype, GlobalLock& lock) { + DTypeRegistry& reg = DTypeRegistry::get(lock); + // Just fill all the names sequentially up to our type ID, + // so that we don't need to deal with recursion. + while (dtype.dtype_id >= reg.names.size()) { + uint32_t next_id = reg.names.size(); + PyPtr name = make_dtype_name(DType{next_id}, lock); + if (!name) return nullptr; + reg.names.push_back(name.release()); + } + return reg.names[dtype.dtype_id]; +} + +int append_to_string_builder(DType dtype, StringBuilder* sb) { + GlobalLock lock; + PyObject* name = dtype_name(dtype, lock); + if (!name) return -1; + sb->append(name); + return 0; +} + + +// ----- DType Python wrapper object ----- + +struct DTypeObject { + DType dtype; + static PyTypeObject pytype; +}; + + +static PyPtr dtype_pyobject_create(DType dtype) { + PyPtr obj = py_create_object(); + if (obj) py_unwrap(obj).dtype = dtype; + return obj; +} + + +PyObject* dtype_to_pyobject(DType dtype, GlobalLock& lock) { + if (!dtype) { + raise(PyExc_ValueError, "Invalid dtype"); + return nullptr; + } + DTypeRegistry& reg = DTypeRegistry::get(lock); + if (dtype.dtype_id >= reg.pyobjects.size()) + reg.pyobjects.resize(dtype.dtype_id + 1); + PyObject*& obj = reg.pyobjects[dtype.dtype_id]; + if (!obj) obj = dtype_pyobject_create(dtype).release(); + return obj; +} + + +static PyObject* DType_reduce(PyObject* self, PyObject*) { + GlobalLock lock; + DType dtype = py_unwrap(self).dtype; + if (dtype.dtype_id < kFirstDerivedDTypeId) { + PyPtr func = getattr(g_module, "_unpickle_basic_dtype"); + if (!func) return nullptr; + PyObject* name = dtype_name(dtype, lock); + if (!name) return nullptr; + return Py_BuildValue("(O(O))", func.get(), name); + } else { + const DerivedDType& derived = get_derived_dtype(dtype, lock); + switch (derived.kind) { + case DerivedDTypeKind::Pointer: + { + PyPtr func = getattr(g_module, "_get_pointer_dtype"); + if (!func) return nullptr; + PyObject* pointee = derived.pointee_dtype + ? dtype_to_pyobject(derived.pointee_dtype, lock) : Py_None; + if (!pointee) return nullptr; + PyObject* space = memory_space_to_pyobject(derived.memory_space, lock); + if (!space) return nullptr; + return Py_BuildValue("(O(OO))", func.get(), pointee, space); + } + case DerivedDTypeKind::ForeignPointer: + { + PyPtr func = getattr(g_module, "_get_foreign_pointer_dtype"); + if (!func) return nullptr; + PyObject* pointee = dtype_to_pyobject(derived.pointee_dtype, lock); + if (!pointee) return nullptr; + return Py_BuildValue("(O(O))", func.get(), pointee); + } + } + CHECK_UNREACHABLE; + } +} + +static PyMethodDef DType_methods[] = { + {"__reduce__", DType_reduce, METH_NOARGS, nullptr}, + {} +}; + +static PyObject* DType_get_name(PyObject* self, void* closure) { + GlobalLock lock; + return Py_NewRef(dtype_name(py_unwrap(self).dtype, lock)); +} + +static PyObject* DType_get_module(PyObject* self, void* closure) { + return PyUnicode_FromString("cuda.tile"); +} + +static PyObject* DType_get_bitwidth(PyObject* self, void* closure) { + GlobalLock lock; + DType dtype = py_unwrap(self).dtype; + return PyLong_FromUnsignedLong(dtype_bitwidth(dtype, lock)); +} + +static PyObject* DType_repr(PyObject* self) { + return to_pyunicode("(self).dtype, "'>").release(); +} + +static PyObject* DType_str(PyObject* self) { + GlobalLock lock; + return Py_NewRef(dtype_name(py_unwrap(self).dtype, lock)); +} + +static PyObject* DType_get_doc(PyObject* self, void* closure) { + GlobalLock lock; + DType dtype = py_unwrap(self).dtype; + if (is_integer_dtype(dtype)) { + StringBuilder sb; + const char* signedness = is_signed_integer_dtype(dtype) ? "signed" : "unsigned"; + + sb.append_many(dtype_bitwidth(dtype, lock), "-bit ", signedness, + " |arithmetic dtype| with values on the interval [", + integer_dtype_min(dtype), ", +", integer_dtype_max(dtype), "]"); + return sb.build().release(); + } + + if (dtype.dtype_id < kFirstDerivedDTypeId) { + const char* doc = basic_dtype_doc( + static_cast<_dtype_detail::BasicDTypeEnum>(dtype.dtype_id)); + if (doc) return PyUnicode_FromString(doc); + } + raise(PyExc_AttributeError, "__doc__"); + return nullptr; +} + +static PyObject *DType_call([[maybe_unused]] PyObject *self, + [[maybe_unused]] PyObject *args, + [[maybe_unused]] PyObject *kwargs) { + raise(PyExc_TypeError, "DType cannot be constructed in pure Python"); + return nullptr; +} + +static PyGetSetDef DType_getsetters[] = { + {"name", DType_get_name, nullptr, "The name of the |data type|"}, + {"__name__", DType_get_name}, + {"__module__", DType_get_module}, + {"bitwidth", DType_get_bitwidth, nullptr, + "The number of bits in an element of the |data type|"}, + {"__doc__", DType_get_doc}, + {} +}; + +PyTypeObject DTypeObject::pytype = { + .tp_name = "cuda.tile.DType", + .tp_basicsize = sizeof(PythonWrapper), + .tp_dealloc = pywrapper_dealloc, + .tp_repr = DType_repr, + .tp_call = DType_call, + .tp_str = DType_str, + .tp_flags = Py_TPFLAGS_DEFAULT, + .tp_doc = + "A *data type* (or *dtype*) describes the type of the objects of an |array|, |tile|, " + "or operation.\n" + "\n" + "|Dtypes| determine how values are stored in memory and how operations on those values are" + "performed.\n", + .tp_methods = DType_methods, + .tp_getset = DType_getsetters, +}; + + +// ----- Python wrappers for is_xxx(dtype) predicates ----- + +using DTypePredicate = bool(DType); +using DTypePredicateWithLock = bool(DType, GlobalLock&); + +template +static PyObject* is_whatever_impl(PyObject* dtype_obj, P&& pred) { + if (!PyObject_TypeCheck(dtype_obj, &DTypeObject::pytype)) { + raise(PyExc_TypeError, "Expected a DType object"); + return nullptr; + } + bool res = pred(py_unwrap(dtype_obj).dtype); + return Py_NewRef(res ? Py_True : Py_False); +} + +template +static PyObject* is_whatever(PyObject* module_self, PyObject* dtype_obj) { + return is_whatever_impl(dtype_obj, P); +} + +template +static PyObject* is_whatever(PyObject* module_self, PyObject* dtype_obj) { + GlobalLock lock; + return is_whatever_impl(dtype_obj, [&lock](DType x){return P(x, lock);}); +} + +static bool is_boolean_dtype(DType dtype) { + return dtype == k_bool_; +} + +// ----- Python wrappers for derived dtype construction ----- + +static PyObject* py_get_pointer_dtype(PyObject* module_self, + PyObject* const* args, Py_ssize_t nargs) { + if (nargs != 2) { + raise(PyExc_TypeError, "Expected 2 positional arguments"); + return nullptr; + } + + PyObject *py_pointee = args[0], *py_memory_space = args[1]; + + DType pointee_dtype; + if (py_pointee == Py_None) { + pointee_dtype = DType::invalid(); + } else if (PyObject_TypeCheck(py_pointee, &DTypeObject::pytype)) { + pointee_dtype = py_unwrap(py_pointee).dtype; + } else { + raise(PyExc_TypeError, "Expected a DType object or None for the first argument, got ", + use_repr(py_pointee)); + return nullptr; + } + + GlobalLock lock; + std::optional space = memory_space_from_pyobject(py_memory_space, lock); + if (!space.has_value()) { + raise(PyExc_TypeError, "Expected a MemorySpace as the second argument, got ", + use_repr(py_memory_space)); + return nullptr; + } + + DType res = pointer_dtype(pointee_dtype, *space, lock); + return Py_NewRef(dtype_to_pyobject(res, lock)); +} + +static PyObject* py_get_foreign_pointer_dtype(PyObject* module_self, PyObject* py_pointee) { + if (!PyObject_TypeCheck(py_pointee, &DTypeObject::pytype)) { + raise(PyExc_TypeError, "Expected a DType object for the first argument, got ", + use_repr(py_pointee)); + return nullptr; + } + DType pointee_dtype = py_unwrap(py_pointee).dtype; + GlobalLock lock; + DType res = foreign_pointer_dtype(pointee_dtype, lock); + return Py_NewRef(dtype_to_pyobject(res, lock)); +} + + +// ----- Python wrappers for dtype queries ----- + +static PyObject* py_integer_dtype_min(PyObject* module_self, PyObject* dtype_obj) { + if (!PyObject_TypeCheck(dtype_obj, &DTypeObject::pytype)) { + raise(PyExc_TypeError, "Expected a DType object, got ", use_repr(dtype_obj)); + return nullptr; + } + DType dtype = py_unwrap(dtype_obj).dtype; + return integer_dtype_min(dtype).to_pylong().release(); +} + +static PyObject* py_integer_dtype_max(PyObject* module_self, PyObject* dtype_obj) { + if (!PyObject_TypeCheck(dtype_obj, &DTypeObject::pytype)) { + raise(PyExc_TypeError, "Expected a DType object, got ", use_repr(dtype_obj)); + return nullptr; + } + DType dtype = py_unwrap(dtype_obj).dtype; + return integer_dtype_max(dtype).to_pylong().release(); +} + +static Result parse_pointer_dtype(PyObject* py_pointer, GlobalLock& lock) { + if (!PyObject_TypeCheck(py_pointer, &DTypeObject::pytype)) + return raise(PyExc_TypeError, "Expected a DType object for the first argument, got ", + use_repr(py_pointer)); + DType pointer_dtype = py_unwrap(py_pointer).dtype; + if (!is_pointer_dtype(pointer_dtype, lock)) + return raise(PyExc_ValueError, pointer_dtype, " is not a pointer dtype"); + return pointer_dtype; +} + +static PyObject* py_pointer_pointee_dtype(PyObject* module_self, PyObject* py_pointer) { + GlobalLock lock; + Result pointer_dtype = parse_pointer_dtype(py_pointer, lock); + if (!pointer_dtype.is_ok()) return nullptr; + DType pointee_dtype = pointer_dtype_pointee(*pointer_dtype, lock); + if (!pointee_dtype) + return Py_NewRef(Py_None); + return Py_NewRef(dtype_to_pyobject(pointee_dtype, lock)); +} + +static PyObject* py_pointer_memory_space(PyObject* module_self, PyObject* py_pointer) { + GlobalLock lock; + Result pointer_dtype = parse_pointer_dtype(py_pointer, lock); + if (!pointer_dtype.is_ok()) return nullptr; + MemorySpace space = pointer_dtype_memory_space(*pointer_dtype, lock); + return Py_NewRef(memory_space_to_pyobject(space, lock)); +} + +static PyObject* py_foreign_pointer_pointee_dtype(PyObject* module_self, PyObject* py_pointer) { + GlobalLock lock; + if (!PyObject_TypeCheck(py_pointer, &DTypeObject::pytype)) { + raise(PyExc_TypeError, "Expected a DType object for the first argument, got ", + use_repr(py_pointer)); + return nullptr; + } + DType foreign_pointer_dtype = py_unwrap(py_pointer).dtype; + if (!is_foreign_pointer_dtype(foreign_pointer_dtype, lock)) { + raise(PyExc_ValueError, foreign_pointer_dtype, " is not a pointer dtype"); + return nullptr; + } + DType pointee_dtype = foreign_pointer_dtype_pointee(foreign_pointer_dtype, lock); + return Py_NewRef(dtype_to_pyobject(pointee_dtype, lock)); +} + + +// ----- Pickle/unpickle helper ----- + +static PyObject* _unpickle_basic_dtype(PyObject* module_self, PyObject* name) { + return PyObject_GetAttr(module_self, name); +} + + +// ----- Module initialization ----- + +static PyMethodDef functions[] = { + {"is_numeric", is_whatever, METH_O, nullptr}, + {"is_boolean", is_whatever, METH_O, nullptr}, + {"is_integral", is_whatever, METH_O, nullptr}, + {"is_signed", is_whatever, METH_O, nullptr}, + {"is_float", is_whatever, METH_O, nullptr}, + {"is_unrestricted_float", is_whatever, METH_O, nullptr}, + {"is_restricted_float", is_whatever, METH_O, nullptr}, + {"is_arithmetic", is_whatever, METH_O, nullptr}, + {"_is_pointer_dtype", is_whatever, METH_O, nullptr}, + {"_is_foreign_pointer_dtype", is_whatever, METH_O, nullptr}, + {"_get_pointer_dtype", reinterpret_cast(py_get_pointer_dtype), + METH_FASTCALL, nullptr}, + {"_get_foreign_pointer_dtype", py_get_foreign_pointer_dtype, METH_O, nullptr}, + {"integer_dtype_min", py_integer_dtype_min, METH_O, nullptr}, + {"integer_dtype_max", py_integer_dtype_max, METH_O, nullptr}, + {"_pointer_pointee_dtype", py_pointer_pointee_dtype, METH_O, nullptr}, + {"_pointer_memory_space", py_pointer_memory_space, METH_O, nullptr}, + {"_foreign_pointer_pointee_dtype", py_foreign_pointer_pointee_dtype, METH_O, nullptr}, + {"_unpickle_basic_dtype", _unpickle_basic_dtype, METH_O, nullptr}, + {} +}; + +#define INIT_STRING_CONSTANT(name, value) \ + if (!(name = PyUnicode_InternFromString(value))) return ErrorRaised + +#define INIT_STRING_IDENT(ident) INIT_STRING_CONSTANT(g_##ident##_pyunicode, #ident) + +Status dtype_init(PyObject* m) { + GlobalLock lock; + + g_module = Py_NewRef(m); + + INIT_STRING_IDENT(pointee_dtype); + INIT_STRING_IDENT(memory_space); + + if (PyType_Ready(&DTypeObject::pytype) < 0) + return ErrorRaised; + + if (PyModule_AddObjectRef(m, "DType", reinterpret_cast(&DTypeObject::pytype)) < 0) + return ErrorRaised; + + if (PyModule_AddFunctions(m, functions) < 0) + return ErrorRaised; + + PyObject* mod_dict = PyModule_GetDict(m); + if (!mod_dict) return ErrorRaised; + + // Add dtypes by name + for (_dtype_detail::BasicDTypeEnum basic_dtype : _dtype_detail::all_basic_dtypes) { + DType dtype = {static_cast(basic_dtype)}; + PyObject* name = dtype_name(dtype, lock); + if (!name) return ErrorRaised; + PyObject* obj = dtype_to_pyobject(dtype, lock); + if (!obj) return ErrorRaised; + if (PyDict_SetItem(mod_dict, name, obj) < 0) + return ErrorRaised; + } + + return OK; +} diff --git a/cext/dtype.h b/cext/dtype.h new file mode 100644 index 00000000..12f3586d --- /dev/null +++ b/cext/dtype.h @@ -0,0 +1,289 @@ +// SPDX-FileCopyrightText: Copyright (c) <2026> NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// +// SPDX-License-Identifier: Apache-2.0 + +#pragma once + +#include "check.h" +#include "py.h" + +#include +#include +#include + +// It may look like this file is overusing X macros, which it probably is. + +class StringBuilder; + +// ----- Memory space ----- + +#define FOREACH_MEMORY_SPACE(X) \ + X(GENERIC, 0, 64) \ + X(GLOBAL, 1, 64) \ + X(SHARED, 3, 32) \ + X(CONSTANT, 4, 64) \ + X(LOCAL, 5, 64) \ + X(TENSOR, 6, 32) \ + X(SHARED_CLUSTER, 7, 32) + +enum class MemorySpace : uint8_t { + #define MEMORY_SPACE_ENUM_ENTRY(name, id, _ptrwidth) \ + name = id, + FOREACH_MEMORY_SPACE(MEMORY_SPACE_ENUM_ENTRY) + #undef MEMORY_SPACE_ENUM_ENTRY +}; + +constexpr uint8_t kMemorySpaceMax = static_cast(MemorySpace::SHARED_CLUSTER); + +const char* memory_space_str(MemorySpace space); + + +static inline uint32_t memory_space_pointer_bitwidth(MemorySpace memory_space) { + switch (memory_space) { + #define MEMORY_SPACE_POINTER_BITWIDTH(name, _id, ptrwidth) \ + case MemorySpace::name: return ptrwidth; + FOREACH_MEMORY_SPACE(MEMORY_SPACE_POINTER_BITWIDTH); + #undef MEMORY_SPACE_POINTER_BITWIDTH + } + CHECK_UNREACHABLE; +} + +// ----- Integer ----- + +struct Integer { + uint64_t bits; + bool is_signed; + + static inline constexpr Integer from_u64(uint64_t x) { + return Integer{x, false}; + } + + static inline constexpr Integer from_i64(int64_t x) { + return Integer{static_cast(x), true}; + } + + PyPtr to_pylong() const { + return steal(is_signed ? PyLong_FromLongLong(static_cast(bits)) + : PyLong_FromUnsignedLongLong(bits)); + } + +}; + +int append_to_string_builder(Integer integer, StringBuilder* sb); + + +// ----- DType ----- + +struct DType { + uint32_t dtype_id; + + static DType invalid() { + return {0}; + } + + explicit operator bool() const { + return dtype_id != 0; + } + + inline bool operator== (DType other) const { + return dtype_id == other.dtype_id; + } + + inline bool operator!= (DType other) const { + return dtype_id != other.dtype_id; + } +}; + + +// (name, bitwdth, signed?, docstring) +#define FOREACH_UNSIGNED_INTEGRAL_DTYPE(X) \ + X(uint8, 8, false, nullptr) \ + X(uint16, 16, false, nullptr) \ + X(uint32, 32, false, nullptr) \ + X(uint64, 64, false, nullptr) + +#define FOREACH_SIGNED_INTEGRAL_DTYPE(X) \ + X(int8, 8, true, nullptr) \ + X(int16, 16, true, nullptr) \ + X(int32, 32, true, nullptr) \ + X(int64, 64, true, nullptr) + +#define FOREACH_UNRESTRICTED_FLOAT_DTYPE(X) \ + X(float16, 16, true, \ + "A IEEE 754 half-precision (16-bit) binary floating-point |arithmetic dtype| " \ + "(see |IEEE 754-2019|).") \ + X(bfloat16, 16, true, \ + "A 16-bit floating-point |arithmetic dtype| with 1 sign bit, 8 exponent bits, " \ + "and 7 mantissa bits.") \ + X(float32, 32, true, \ + "A IEEE 754 single-precision (32-bit) binary floating-point |arithmetic dtype| " \ + "(see |IEEE 754-2019|).") \ + X(float64, 64, true, \ + "A IEEE 754 double-precision (64-bit) binary floating-point |arithmetic dtype| " \ + "(see |IEEE 754-2019|).") + +#define FOREACH_RESTRICTED_FLOAT_DTYPE(X) \ + X(tfloat32, 32, true, \ + "A 32-bit tensor floating-point |numeric dtype| with 1 sign bit, 8 exponent bits, " \ + "and 10 mantissa bits (19-bit representation stored in 32-bit container).") \ + X(float8_e4m3fn, 8, true, \ + "An 8-bit floating-point |numeric dtype| with 1 sign bit, " \ + "4 exponent bits, and 3 mantissa bits.") \ + X(float8_e5m3fnu, 8, false, \ + "An 8-bit floating-point |numeric dtype| with no sign bit, " \ + "5 exponent bits, and 3 mantissa bits.") \ + X(float8_e5m2, 8, true, \ + "An 8-bit floating-point |numeric dtype| with 1 sign bit, " \ + "5 exponent bits, and 2 mantissa bits.") \ + X(float8_e8m0fnu, 8, false, \ + "An 8-bit floating-point |numeric dtype| with no sign bit, " \ + "8 exponent bits, and 0 mantissa bits.") \ + X(float4_e2m1fn, 4, true, \ + "A 4-bit floating-point |numeric dtype| with 1 sign bit, " \ + "2 exponent bits, and 1 mantissa bit.") \ + X(float6_e2m3fn, 6, true, \ + "A 6-bit floating-point numeric dtype with 1 sign bit, " \ + "2 exponent bits, and 3 mantissa bits.") \ + X(float6_e3m2fn, 6, true, \ + "A 6-bit floating-point numeric dtype with 1 sign bit, " \ + "3 exponent bits, and 2 mantissa bits.") + + +#define FOREACH_NUMERIC_DTYPE(X) \ + X(bool_, 8, false, "An 8-bit boolean |arithmetic dtype|.") \ + FOREACH_UNSIGNED_INTEGRAL_DTYPE(X) \ + FOREACH_SIGNED_INTEGRAL_DTYPE(X) \ + FOREACH_UNRESTRICTED_FLOAT_DTYPE(X) \ + FOREACH_RESTRICTED_FLOAT_DTYPE(X) + +#define FOREACH_BASIC_DTYPE(X) \ + FOREACH_NUMERIC_DTYPE(X) \ + X(mbarrier, 64, false, "An opaque dtype representing an mbarrier state.") \ + X(cluster_launch_control_token, 128, false, \ + "An opaque dtype representing a cluster launch control token value.") \ + X(tensor_map_descriptor, 8 * sizeof(CUtensorMap), false, \ + "An opaque dtype representing a tensor map descriptor.") + + +namespace _dtype_detail { + enum class BasicDTypeEnum : uint32_t { + Invalid = 0, + #define BASIC_DTYPE_ENUM_ENTRY(name, _bitwidth, _signed, _doc) \ + name, + FOREACH_BASIC_DTYPE(BASIC_DTYPE_ENUM_ENTRY) + #undef BASIC_DTYPE_ENUM_ENTRY + }; + + // Define constants like `unsigned_int_first` and `unsigned_int_last`, + // for quickly checking a dtype category. + #define BASIC_DTYPE_ENUM(name, _bitwidth, _signed, _doc) \ + BasicDTypeEnum::name, + + #define BASIC_DTYPE_FIRST_LAST(category, xmacro) \ + static constexpr BasicDTypeEnum all_##category##_dtypes[] = { \ + xmacro(BASIC_DTYPE_ENUM) \ + }; \ + static constexpr uint32_t category##_first = static_cast( \ + all_##category##_dtypes[0]); \ + static constexpr uint32_t category##_last = static_cast( \ + all_##category##_dtypes[std::extent_v - 1]); + + BASIC_DTYPE_FIRST_LAST(unsigned_int, FOREACH_UNSIGNED_INTEGRAL_DTYPE); + BASIC_DTYPE_FIRST_LAST(signed_int, FOREACH_SIGNED_INTEGRAL_DTYPE); + BASIC_DTYPE_FIRST_LAST(unrestricted_float, FOREACH_UNRESTRICTED_FLOAT_DTYPE); + BASIC_DTYPE_FIRST_LAST(restricted_float, FOREACH_RESTRICTED_FLOAT_DTYPE); + BASIC_DTYPE_FIRST_LAST(numeric, FOREACH_NUMERIC_DTYPE); + BASIC_DTYPE_FIRST_LAST(basic, FOREACH_BASIC_DTYPE); + + #undef BASIC_DTYPE_ENUM + #undef BASIC_DTYPE_FIRST_LAST + + extern const uint32_t basic_dtype_bitwidth[basic_last + 1]; + extern const bool basic_dtype_is_signed[basic_last + 1]; + + bool is_derived_dtype_pointer(DType dtype, GlobalLock& lock); + bool is_derived_dtype_foreign_pointer(DType dtype, GlobalLock& lock); + uint32_t derived_dtype_bitwidth(DType dtype, GlobalLock& lock); +} // namespace _dtype_detail + +// Define constants like +// static constexpr DType k_int32 = ...; +#define BASIC_DTYPE_CONSTANT(name, _bitwidth, _signed, _doc) \ + static constexpr DType k_##name = {static_cast(_dtype_detail::BasicDTypeEnum::name)}; +FOREACH_BASIC_DTYPE(BASIC_DTYPE_CONSTANT) +#undef BASIC_DTYPE_CONSTANT + +static constexpr uint32_t kFirstDerivedDTypeId = _dtype_detail::basic_last + 1; + +static inline bool is_unsigned_integer_dtype(DType dtype) { + return dtype.dtype_id >= _dtype_detail::unsigned_int_first + && dtype.dtype_id <= _dtype_detail::unsigned_int_last; +} + +static inline bool is_signed_integer_dtype(DType dtype) { + return dtype.dtype_id >= _dtype_detail::signed_int_first + && dtype.dtype_id <= _dtype_detail::signed_int_last; +} + +static inline bool is_integer_dtype(DType dtype) { + return is_unsigned_integer_dtype(dtype) || is_signed_integer_dtype(dtype); +} + +Integer integer_dtype_min(DType dtype); + +Integer integer_dtype_max(DType dtype); + +static inline bool is_unrestricted_float_dtype(DType dtype) { + return dtype.dtype_id >= _dtype_detail::unrestricted_float_first + && dtype.dtype_id <= _dtype_detail::unrestricted_float_last; +} + +static inline bool is_restricted_float_dtype(DType dtype) { + return dtype.dtype_id >= _dtype_detail::restricted_float_first + && dtype.dtype_id <= _dtype_detail::restricted_float_last; +} + +static inline bool is_float_dtype(DType dtype) { + return is_unrestricted_float_dtype(dtype) || is_restricted_float_dtype(dtype); +} + +static inline bool is_arithmetic_numeric_dtype(DType dtype) { + return dtype == k_bool_ || is_integer_dtype(dtype) || is_unrestricted_float_dtype(dtype); +} + +static inline bool is_pointer_dtype(DType dtype, GlobalLock& lock) { + return dtype.dtype_id >= kFirstDerivedDTypeId + && _dtype_detail::is_derived_dtype_pointer(dtype, lock); +} + +static inline bool is_foreign_pointer_dtype(DType dtype, GlobalLock& lock) { + return dtype.dtype_id >= kFirstDerivedDTypeId + && _dtype_detail::is_derived_dtype_foreign_pointer(dtype, lock); +} + +static inline bool is_numeric_dtype(DType dtype) { + return dtype.dtype_id >= _dtype_detail::numeric_first + && dtype.dtype_id <= _dtype_detail::numeric_last; +} + +static inline bool is_signed_numeric_dtype(DType dtype) { + return is_numeric_dtype(dtype) && _dtype_detail::basic_dtype_is_signed[dtype.dtype_id]; +} + +DType pointer_dtype_pointee(DType pointer_dtype, GlobalLock& lock); + +DType foreign_pointer_dtype_pointee(DType foreign_pointer_dtype, GlobalLock& lock); + +MemorySpace pointer_dtype_memory_space(DType pointer_dtype, GlobalLock& lock); + +static inline uint32_t dtype_bitwidth(DType dtype, GlobalLock& lock) { + CHECK(dtype); + if (dtype.dtype_id < kFirstDerivedDTypeId) + return _dtype_detail::basic_dtype_bitwidth[dtype.dtype_id]; + + return _dtype_detail::derived_dtype_bitwidth(dtype, lock); +} + +int append_to_string_builder(DType dtype, StringBuilder* sb); + +Status dtype_init(PyObject* m); diff --git a/cext/module.cpp b/cext/module.cpp index 38885634..4502b500 100644 --- a/cext/module.cpp +++ b/cext/module.cpp @@ -5,6 +5,7 @@ #include "py.h" #include "compiled_host.h" +#include "dtype.h" #include "tile_kernel.h" #include "cuda_helper.h" #include "coroutine_util.h" @@ -34,6 +35,9 @@ PyMODINIT_FUNC PyInit__cext() { return nullptr; #endif + if (!dtype_init(m.get())) + return nullptr; + if (!tile_kernel_init(m.get())) return nullptr; diff --git a/cext/py.h b/cext/py.h index 298a7b6f..1175ae18 100644 --- a/cext/py.h +++ b/cext/py.h @@ -10,7 +10,9 @@ #include "ref_ptr.h" #include "vec.h" #include +#include #include +#include using PyPtr = RefPtr; @@ -96,13 +98,23 @@ T& py_unwrap(PyObject* pyobj) { } template -PyObject* pywrapper_new(PyTypeObject* type, PyObject*, PyObject*) { +T& py_unwrap(const PyPtr& pyobj) { + return py_unwrap(pyobj.get()); +} + +template +PyPtr py_create_object(PyTypeObject* type = &T::pytype) { PyObject* ret = type->tp_alloc(type, 0); - if (!ret) return nullptr; + if (!ret) return {}; T& obj = py_unwrap(ret); new (&obj) T(); - return ret; + return steal(ret); +} + +template +PyObject* pywrapper_new(PyTypeObject* type, PyObject*, PyObject*) { + return py_create_object(type).release(); } template @@ -303,93 +315,41 @@ class StringBuilder : StringBuilderImpl { discard(); } - void append(const char* s) { + void append_cstring(const char* s) { handle_error([=] { return write_ascii(s); }); } - void append(char c) { + void append_char(char c) { Py_UCS4 usc4 = static_cast(c); handle_error([=] { return write_char(usc4); }); } - void append(unsigned int x) { - append_sprintf<30>("%u", x); - } - - void append(int x) { - append_sprintf<30>("%d", x); - } - - void append(unsigned long x) { - append_sprintf<30>("%lu", x); - } - - void append(long x) { - append_sprintf<30>("%ld", x); - } - - void append(unsigned long long x) { - append_sprintf<30>("%llu", x); - } - - void append(long long x) { - append_sprintf<30>("%lld", x); - } - - void append(PyObject* obj) { - handle_error([=] { return obj ? write_str(obj) : write_ascii("(null)"); }); - } - - void append(UseRepr u) { - handle_error([=] { return u.obj ? write_repr(u.obj) : write_ascii("(null)"); }); - } - - void append(const PyPtr& obj) { - append(obj.get()); - } - - template >> - void append(const T& value) { - append(static_cast>(value)); + void append_pyobject_str(PyObject* obj) { + handle_error([=] { return write_str(obj); }); } - template - void append(const T* ptr) { - append_sprintf<30>("%p", ptr); + void append_pyobject_repr(PyObject* obj) { + handle_error([=] { return write_repr(obj); }); } - template - void append(const Vec& vec) { - append("["); - const char* comma = ""; - for (const T& x : vec) { - append(comma); - append(x); - comma = ", "; - } - append("]"); - } - - template - void append(const Result& res) { - if (res.is_ok()) { - append("OK("); - append(*res); - append(")"); - } else { - append("ErrorRaised"); - } + template + void append_sprintf(const char* fmt, Args&&... args) { + handle_error([&] { + char buf[BufSize]; + int r = PyOS_snprintf(buf, sizeof buf, fmt, std::forward(args)...); + if (r < 0 || r >= (int) sizeof buf) { + PyErr_SetString(PyExc_RuntimeError, "snprintf() failed"); + return -1; + } + return write_ascii(buf, r); + }); } template - void append(const std::optional& opt) { - if (opt.has_value()) { - append("std::optional{"); - append(*opt); - append("}"); - } else { - append("std::nullopt"); - } + void append(T&& val) { + handle_error([&] { + return append_to_string_builder(std::forward(val), this); + }); } template @@ -397,6 +357,7 @@ class StringBuilder : StringBuilderImpl { (append(std::forward(args)), ...); } + PyPtr build() { if (error_) return {}; return steal(finish()); @@ -414,19 +375,6 @@ class StringBuilder : StringBuilderImpl { private: bool error_; - template - void append_sprintf(const char* fmt, Args&&... args) { - handle_error([&] { - char buf[BufSize]; - int r = PyOS_snprintf(buf, sizeof buf, fmt, std::forward(args)...); - if (r < 0 || r >= (int) sizeof buf) { - PyErr_SetString(PyExc_RuntimeError, "snprintf() failed"); - return -1; - } - return write_ascii(buf, r); - }); - } - template void handle_error(F&& func) { if (!error_) { @@ -436,6 +384,113 @@ class StringBuilder : StringBuilderImpl { } }; +static inline int append_to_string_builder(const char* s, StringBuilder* sb) { + sb->append_cstring(s); + return 0; +} + +static inline int append_to_string_builder(char c, StringBuilder* sb) { + sb->append_char(c); + return 0; +} + +static inline int append_to_string_builder(unsigned int x, StringBuilder* sb) { + sb->append_sprintf<30>("%u", x); + return 0; +} + +static inline int append_to_string_builder(int x, StringBuilder* sb) { + sb->append_sprintf<30>("%d", x); + return 0; +} + +static inline int append_to_string_builder(unsigned long x, StringBuilder* sb) { + sb->append_sprintf<30>("%lu", x); + return 0; +} + +static inline int append_to_string_builder(long x, StringBuilder* sb) { + sb->append_sprintf<30>("%ld", x); + return 0; +} + +static inline int append_to_string_builder(unsigned long long x, StringBuilder* sb) { + sb->append_sprintf<30>("%llu", x); + return 0; +} + +static inline int append_to_string_builder(long long x, StringBuilder* sb) { + sb->append_sprintf<30>("%lld", x); + return 0; +} + +static inline int append_to_string_builder(PyObject* obj, StringBuilder* sb) { + if (obj) sb->append_pyobject_str(obj); + else sb->append("(null)"); + return 0; +} + +static inline int append_to_string_builder(UseRepr u, StringBuilder* sb) { + if (u.obj) sb->append_pyobject_repr(u.obj); + else sb->append("(null)"); + return 0; +} + +static inline int append_to_string_builder(const PyPtr& obj, StringBuilder* sb) { + sb->append(obj.get()); + return 0; +} + +template +static inline int append_to_string_builder(const T* ptr, StringBuilder* sb) { + sb->append_sprintf<30>("%p", ptr); + return 0; +} + +template +static inline int append_to_string_builder(const Vec& vec, StringBuilder* sb) { + sb->append("["); + const char* comma = ""; + for (const T& x : vec) { + sb->append(comma); + sb->append(x); + comma = ", "; + } + sb->append("]"); + return 0; +} + +template +static inline int append_to_string_builder(const Result& res, StringBuilder* sb) { + if (res.is_ok()) { + sb->append("OK("); + sb->append(*res); + sb->append(")"); + } else { + sb->append("ErrorRaised"); + } + return 0; +} + +template +static inline int append_to_string_builder(const std::optional& opt, StringBuilder* sb) { + if (opt.has_value()) { + sb->append("std::optional{"); + sb->append(*opt); + sb->append("}"); + } else { + sb->append("std::nullopt"); + } + return 0; +} + +template >> +static inline int append_to_string_builder(const T& value, StringBuilder* sb) { + sb->append(static_cast>(value)); + return 0; +} + + template ErrorRaised_t raise(PyObject* exctype, Args&&... message) { StringBuilder builder; diff --git a/experimental/cuda-lang/src/cuda/lang/_datatype.py b/experimental/cuda-lang/src/cuda/lang/_datatype.py index eb853529..9abf5861 100644 --- a/experimental/cuda-lang/src/cuda/lang/_datatype.py +++ b/experimental/cuda-lang/src/cuda/lang/_datatype.py @@ -7,7 +7,6 @@ from cuda.tile._memory_model import MemorySpace from cuda.tile._datatype import ( DType, - NumericDTypeCategory, bfloat16, bool_, float8_e8m0fnu, @@ -42,27 +41,14 @@ is_pointer_dtype, pointer_dtype, opaque_pointer_dtype, - _define_dtype, - _DTypeDefinition, PointerInfo, numeric_dtype_category, ) -from cuda.tile import _cext - -# Lang-specific types. -float6_e2m3fn = _define_dtype( - "float6_e2m3fn", - _DTypeDefinition(bitwidth=6, numeric_category=NumericDTypeCategory.RestrictedFloat), +from cuda.tile._cext import ( + mbarrier, cluster_launch_control_token, float6_e2m3fn, float6_e3m2fn, + tensor_map_descriptor, ) -float6_e2m3fn.__doc__ = """A 6-bit floating-point numeric dtype with 1 sign bit, -2 exponent bits, and 3 mantissa bits.""" -float6_e3m2fn = _define_dtype( - "float6_e3m2fn", - _DTypeDefinition(bitwidth=6, numeric_category=NumericDTypeCategory.RestrictedFloat), -) -float6_e3m2fn.__doc__ = """A 6-bit floating-point numeric dtype with 1 sign bit, -3 exponent bits, and 2 mantissa bits.""" arithmetic_float_dtypes = ( float16, @@ -81,15 +67,6 @@ float6_e3m2fn, ) -mbarrier = _define_dtype('mbarrier', _DTypeDefinition(bitwidth=64)) -cluster_launch_control_token = _define_dtype( - "cluster_launch_control_token", _DTypeDefinition(bitwidth=128) -) -tensor_map_descriptor = _define_dtype( - "tensor_map_descriptor", - _DTypeDefinition(bitwidth=8 * _cext._TENSOR_MAP_DESCRIPTOR_BYTES), -) - def to_torch_dtype(dtype: DType, /): if not isinstance(dtype, DType): diff --git a/experimental/cuda-lang/test/test_dtype.py b/experimental/cuda-lang/test/test_dtype.py index d81f9834..414bb7d7 100644 --- a/experimental/cuda-lang/test/test_dtype.py +++ b/experimental/cuda-lang/test/test_dtype.py @@ -2,14 +2,32 @@ # # SPDX-License-Identifier: Apache-2.0 +import pickle + import torch.cuda import pytest import cuda.lang as cl import cuda.lang._datatype as datatype +from cuda.tile import _cext from test.util import compile_kernel +def test_tensor_map_descriptor_dtype(): + dtype = cl.tensor_map_descriptor + assert isinstance(dtype, datatype.DType) + assert dtype is _cext.tensor_map_descriptor + assert dtype.name == "tensor_map_descriptor" + assert dtype.bitwidth == 8 * _cext._TENSOR_MAP_DESCRIPTOR_BYTES + assert not datatype.is_numeric(dtype) + assert not cl.is_pointer_dtype(dtype) + assert pickle.loads(pickle.dumps(dtype)) is dtype + + pointer_dtype = cl.pointer_dtype(dtype) + assert cl.PointerInfo(pointer_dtype).pointee_dtype is dtype + assert pickle.loads(pickle.dumps(pointer_dtype)) is pointer_dtype + + @pytest.mark.parametrize("dtype", datatype.arithmetic_float_dtypes) def test_arithmetic_float_dtype(dtype): assert datatype.is_numeric(dtype) diff --git a/src/cuda/tile/_cext.pyi b/src/cuda/tile/_cext.pyi index 4caf1edf..462eadc1 100644 --- a/src/cuda/tile/_cext.pyi +++ b/src/cuda/tile/_cext.pyi @@ -5,11 +5,78 @@ import enum from typing import Any, Sequence, TypeAlias from cuda.tile._context import TileContextConfig - +from cuda.tile._memory_model import MemorySpace Dim3: TypeAlias = tuple[int] | tuple[int, int] | tuple[int, int, int] +class DType: + """A *data type* (or *dtype*) describes the type of the objects of an |array|, |tile|, or + operation. + + |Dtypes| determine how values are stored in memory and how operations on those values are + performed. + """ + + @property + def bitwidth(self): + """The number of bits in an element of the |data type|.""" + + @property + def name(self): + """The name of the |data type|.""" + + def __call__(self, value, /): + """Construct a Scalar of this |data type| from a value.""" + + +def is_numeric(t: DType, /) -> bool: ... +def is_boolean(t: DType, /) -> bool: ... +def is_integral(t: DType, /) -> bool: ... +def is_signed(t: DType, /) -> bool: ... +def is_float(t: DType, /) -> bool: ... +def is_unrestricted_float(t: DType, /) -> bool: ... +def is_restricted_float(t: DType, /) -> bool: ... +def is_arithmetic(t: DType, /) -> bool: ... +def _is_pointer_dtype(t: DType, /) -> bool: ... +def _is_foreign_pointer_dtype(t: DType, /) -> bool: ... + +def _get_pointer_dtype(pointee_dtype: DType | None, memory_space: MemorySpace, /) -> DType: ... +def _get_foreign_pointer_dtype(pointee_dtype: DType, /) -> DType: ... +def _pointer_pointee_dtype(pointer_dtype: DType, /) -> DType: ... +def _pointer_memory_space(pointer_dtype: DType, /) -> MemorySpace: ... +def _foreign_pointer_pointee_dtype(foreign_pointer_dtype: DType, /) -> DType: ... + +def integer_dtype_min(dtype: DType) -> int: ... +def integer_dtype_max(dtype: DType) -> int: ... + + +bool_: DType +uint8: DType +uint16: DType +uint32: DType +uint64: DType +int8: DType +int16: DType +int32: DType +int64: DType +float16: DType +float32: DType +float64: DType +bfloat16: DType +tfloat32: DType +float8_e4m3fn: DType +float8_e5m2: DType +float8_e8m0fnu: DType +float8_e5m3fnu: DType +float4_e2m1fn: DType +mbarrier: DType +cluster_launch_control_token: DType +tensor_map_descriptor: DType +float6_e2m3fn: DType +float6_e3m2fn: DType + + def launch(stream, grid: Dim3, kernel, diff --git a/src/cuda/tile/_datatype.py b/src/cuda/tile/_datatype.py index 356ec865..5f3d4104 100644 --- a/src/cuda/tile/_datatype.py +++ b/src/cuda/tile/_datatype.py @@ -4,16 +4,23 @@ from __future__ import annotations -import threading -from dataclasses import dataclass -from functools import cache from typing import Optional, Tuple from enum import IntEnum from cuda.tile._exception import TileTypeError -from cuda.tile._execution import function, stub +from cuda.tile._execution import stub from cuda.tile._memory_model import MemorySpace import cuda.tile._bytecode as bc +from cuda.tile._cext import ( + DType, is_numeric, is_boolean, is_integral, is_signed, is_float, is_unrestricted_float, + is_restricted_float, is_arithmetic, _is_pointer_dtype, _is_foreign_pointer_dtype, + _pointer_pointee_dtype, _pointer_memory_space, _foreign_pointer_pointee_dtype, + _get_pointer_dtype, _get_foreign_pointer_dtype, + integer_dtype_min, integer_dtype_max, + bool_, uint8, uint16, uint32, uint64, int8, int16, int32, int64, + float16, bfloat16, float32, float64, tfloat32, float8_e4m3fn, float8_e5m2, float8_e8m0fnu, + float8_e5m3fnu, float4_e2m1fn +) __all__ = ["bool_", "uint8", "uint16", "uint32", "uint64", @@ -21,50 +28,8 @@ "float16", "float32", "float64", "bfloat16", "tfloat32", "float8_e4m3fn", "float8_e5m2", "float8_e8m0fnu", "float8_e5m3fnu", "float4_e2m1fn", "DType", - "foreign_pointer_dtype"] - - -class DType: - """A *data type* (or *dtype*) describes the type of the objects of an |array|, |tile|, or - operation. - - |Dtypes| determine how values are stored in memory and how operations on those values are - performed. - |Dtypes| are immutable. - - |Dtypes| can be used in |host code| and |tile code|. - They can be |kernel| parameters. - """ - - def __new__(self): - raise TypeError("DType objects cannot be created") - - def __reduce__(self): - return _define_dtype, (self.__name__, _dtype_defs[self]) - - @property - @function(host=True, tile=False) - def bitwidth(self): - """The number of bits in an element of the |data type|.""" - return _dtype_defs[self].bitwidth - - @property - @function(host=True, tile=False) - def name(self): - """The name of the |data type|.""" - return self.__name__ - - @function(host=True, tile=False) - def __repr__(self): - return f"" - - @function(host=True, tile=False) - def __str__(self): - return self.__name__ - - @stub - def __call__(self, value, /): - """Construct a Scalar of this |data type| from a value.""" + "foreign_pointer_dtype", "is_numeric", "is_boolean", "is_integral", "is_signed", + "is_arithmetic", "is_float", "is_unrestricted_float", "is_restricted_float"] class NumericDTypeCategory(IntEnum): @@ -82,15 +47,6 @@ def pytype(self) -> type: case NumericDTypeCategory.RestrictedFloat: return float case _: assert False, self - @property - def arithmetic(self) -> bool: - match self: - case NumericDTypeCategory.Boolean: return True - case NumericDTypeCategory.Integral: return True - case NumericDTypeCategory.Float: return True - case NumericDTypeCategory.RestrictedFloat: return False - case _: assert False, self - class IntegerInfo: """ @@ -98,11 +54,9 @@ class IntegerInfo: """ @stub(host=True) def __init__(self, dtype: DType): - definition = _dtype_defs[dtype] - if not isinstance(definition, _IntegerDTypeDefinition): + if not is_integral(dtype): raise TypeError(f"'{dtype}' is not an integer dtype") self._dtype = dtype - self._definition = definition @property def dtype(self) -> DType: @@ -110,15 +64,15 @@ def dtype(self) -> DType: @property def bits(self) -> int: - return self._definition.bitwidth + return self._dtype.bitwidth @property def min(self) -> int: - return self._definition.get_min_value() + return integer_dtype_min(self._dtype) @property def max(self) -> int: - return self._definition.get_max_value() + return integer_dtype_max(self._dtype) def __eq__(self, other): return isinstance(other, IntegerInfo) and self._dtype == other._dtype @@ -127,136 +81,6 @@ def __hash__(self): return hash(self._dtype) -@dataclass(frozen=True, kw_only=True) -class _DTypeDefinition: - bitwidth: int - numeric_category: NumericDTypeCategory | None = None - simple_bytecode_type: bc.SimpleType | None = None - - -@dataclass(frozen=True, kw_only=True) -class _IntegerDTypeDefinition(_DTypeDefinition): - signed: bool - - def get_min_value(self) -> int: - return -(1 << (self.bitwidth - 1)) if self.signed else 0 - - def get_max_value(self) -> int: - return (1 << (self.bitwidth - 1)) - 1 if self.signed else (1 << self.bitwidth) - 1 - - -@dataclass(frozen=True, kw_only=True) -class _PointerDTypeDefinition(_DTypeDefinition): - pointee_dtype: DType | None # None for opaque pointers - memory_space: MemorySpace - - -@dataclass(frozen=True, kw_only=True) -class _ForeignPointerDTypeDefinition(_DTypeDefinition): - pointee_dtype: DType - - -_dtype_defs: dict[DType, _DTypeDefinition] = dict() -_dtype_by_name: dict[str, DType] = dict() -_dtype_lock = threading.Lock() - - -def _define_dtype(name: str, definition: _DTypeDefinition) -> DType: - assert isinstance(definition, _DTypeDefinition) - with _dtype_lock: - if name in _dtype_by_name: - existing = _dtype_by_name[name] - assert _dtype_defs[existing] == definition - return existing - - dtype = object.__new__(DType) - dtype.__name__ = name - _dtype_defs[dtype] = definition - _dtype_by_name[name] = dtype - return dtype - - -def _numeric_dtype(name: str, - bitwidth: int, - category: NumericDTypeCategory, - bc_type: bc.SimpleType) -> DType: - definition = _DTypeDefinition(bitwidth=bitwidth, - numeric_category=category, - simple_bytecode_type=bc_type) - return _define_dtype(name, definition) - - -def _integer_dtype(name: str, bitwidth: int, signed: bool, bc_type: bc.SimpleType) -> DType: - definition = _IntegerDTypeDefinition(bitwidth=bitwidth, - numeric_category=NumericDTypeCategory.Integral, - simple_bytecode_type=bc_type, - signed=signed) - dtype = _define_dtype(name, definition) - signedness = "signed" if signed else "unsigned" - dtype.__doc__ = (f"{bitwidth}-bit {signedness} integer |arithmetic dtype| with values" - f" on the interval" - f" [{definition.get_min_value()}, +{definition.get_max_value()}]") - return dtype - - -bool_ = _numeric_dtype('bool_', 8, NumericDTypeCategory.Boolean, bc.SimpleType.I1) -bool_.__doc__ = """A 8-bit |arithmetic dtype| (``True`` or ``False``).""" - -uint8 = _integer_dtype('uint8', 8, False, bc.SimpleType.I8) -uint16 = _integer_dtype('uint16', 16, False, bc.SimpleType.I16) -uint32 = _integer_dtype('uint32', 32, False, bc.SimpleType.I32) -uint64 = _integer_dtype('uint64', 64, False, bc.SimpleType.I64) -int8 = _integer_dtype('int8', 8, True, bc.SimpleType.I8) -int16 = _integer_dtype('int16', 16, True, bc.SimpleType.I16) -int32 = _integer_dtype('int32', 32, True, bc.SimpleType.I32) -int64 = _integer_dtype('int64', 64, True, bc.SimpleType.I64) - -float16 = _numeric_dtype('float16', 16, NumericDTypeCategory.Float, bc.SimpleType.F16) -float16.__doc__ = """A IEEE 754 half-precision (16-bit) binary floating-point |arithmetic dtype| \ -(see |IEEE 754-2019|).""" - -float32 = _numeric_dtype('float32', 32, NumericDTypeCategory.Float, bc.SimpleType.F32) -float32.__doc__ = """A IEEE 754 single-precision (32-bit) binary floating-point |arithmetic dtype| \ -(see |IEEE 754-2019|).""" - -float64 = _numeric_dtype('float64', 64, NumericDTypeCategory.Float, bc.SimpleType.F64) -float64.__doc__ = """A IEEE 754 double-precision (64-bit) binary floating-point |arithmetic dtype| \ -(see |IEEE 754-2019|).""" - -bfloat16 = _numeric_dtype('bfloat16', 16, NumericDTypeCategory.Float, bc.SimpleType.BF16) -bfloat16.__doc__ = """A 16-bit floating-point |arithmetic dtype| with 1 sign bit, 8 exponent bits, \ -and 7 mantissa bits.""" - -tfloat32 = _numeric_dtype("tfloat32", 32, NumericDTypeCategory.RestrictedFloat, bc.SimpleType.TF32) -tfloat32.__doc__ = """A 32-bit tensor floating-point |numeric dtype| with 1 sign \ -bit, 8 exponent bits, and 10 mantissa bits (19-bit representation stored in 32-bit container).""" - -float8_e4m3fn = _numeric_dtype("float8_e4m3fn", 8, NumericDTypeCategory.RestrictedFloat, - bc.SimpleType.F8E4M3FN) -float8_e4m3fn.__doc__ = """An 8-bit floating-point |numeric dtype| with 1 sign bit, \ -4 exponent bits, and 3 mantissa bits.""" - -float8_e5m3fnu = _numeric_dtype("float8_e5m3fnu", 8, NumericDTypeCategory.RestrictedFloat, - bc.SimpleType.FNV8E5M3FNU) -float8_e5m3fnu.__doc__ = """An 8-bit floating-point |numeric dtype| with no sign bit, \ -5 exponent bits, and 3 mantissa bits.""" - -float8_e5m2 = _numeric_dtype("float8_e5m2", 8, NumericDTypeCategory.RestrictedFloat, - bc.SimpleType.F8E5M2) -float8_e5m2.__doc__ = """An 8-bit floating-point |numeric dtype| with 1 sign bit, \ -5 exponent bits, and 2 mantissa bits.""" - -float8_e8m0fnu = _numeric_dtype("float8_e8m0fnu", 8, NumericDTypeCategory.RestrictedFloat, - bc.SimpleType.F8E8M0FNU) -float8_e8m0fnu.__doc__ = """An 8-bit floating-point |numeric dtype| with no sign bit, \ -8 exponent bits, and 0 mantissa bits.""" - -float4_e2m1fn = _numeric_dtype("float4_e2m1fn", 4, NumericDTypeCategory.RestrictedFloat, - bc.SimpleType.F4E2M1FN) -float4_e2m1fn.__doc__ = """A 4-bit floating-point |numeric dtype| with 1 sign bit, \ -2 exponent bits, and 1 mantissa bit.""" - - default_int_type = int32 default_float_type = float32 @@ -268,44 +92,45 @@ def _integer_dtype(name: str, bitwidth: int, signed: bool, bc_type: bc.SimpleTyp signed_integral_dtypes = [int64, int32, int16, int8] -def is_numeric(t: DType) -> bool: - return _dtype_defs[t].numeric_category is not None - - def numeric_dtype_category(t: DType) -> NumericDTypeCategory: - cat = _dtype_defs[t].numeric_category - if cat is None: + if is_boolean(t): + return NumericDTypeCategory.Boolean + elif is_integral(t): + return NumericDTypeCategory.Integral + elif is_unrestricted_float(t): + return NumericDTypeCategory.Float + elif is_restricted_float(t): + return NumericDTypeCategory.RestrictedFloat + else: + assert not is_numeric(t) raise ValueError(f"{t} is not a numeric dtype") - return cat - - -def dtype_simple_bytecode_type(t: DType) -> bc.SimpleType: - ret = _dtype_defs[t].simple_bytecode_type - assert ret is not None - return ret - - -def is_boolean(t: DType) -> bool: - return _dtype_defs[t].numeric_category == NumericDTypeCategory.Boolean -def is_integral(t: DType) -> bool: - return _dtype_defs[t].numeric_category == NumericDTypeCategory.Integral +_dtype_to_simple_bytecode_type = { + bool_: bc.SimpleType.I1, + uint8: bc.SimpleType.I8, + uint16: bc.SimpleType.I16, + uint32: bc.SimpleType.I32, + uint64: bc.SimpleType.I64, + int8: bc.SimpleType.I8, + int16: bc.SimpleType.I16, + int32: bc.SimpleType.I32, + int64: bc.SimpleType.I64, + float16: bc.SimpleType.F16, + bfloat16: bc.SimpleType.BF16, + float32: bc.SimpleType.F32, + tfloat32: bc.SimpleType.TF32, + float64: bc.SimpleType.F64, + float8_e4m3fn: bc.SimpleType.F8E4M3FN, + float8_e5m2: bc.SimpleType.F8E5M2, + float8_e8m0fnu: bc.SimpleType.F8E8M0FNU, + float4_e2m1fn: bc.SimpleType.F4E2M1FN, + float8_e5m3fnu: bc.SimpleType.FNV8E5M3FNU, +} -def is_signed(t: DType) -> bool: - """Returns True if the |dtype| is a signed numeric type, such as a signed integer or - a floating-point type.""" - info = _dtype_defs[t] - match info.numeric_category: - case None: return False - case NumericDTypeCategory.Boolean: return False - case NumericDTypeCategory.Integral: - assert isinstance(info, _IntegerDTypeDefinition) - return info.signed - case NumericDTypeCategory.Float: return True - case NumericDTypeCategory.RestrictedFloat: return True - case _: assert False, info.numeric_category +def dtype_simple_bytecode_type(t: DType) -> bc.SimpleType: + return _dtype_to_simple_bytecode_type[t] def integer_dtype(bitwidth: int, *, signed: bool) -> DType: @@ -329,26 +154,6 @@ def get_signedness(t: DType) -> bc.Signedness: return _signedness[is_signed(t)] -def is_float(t: DType) -> bool: - return _dtype_defs[t].numeric_category in (NumericDTypeCategory.Float, - NumericDTypeCategory.RestrictedFloat) - - -def is_unrestricted_float(t: DType) -> bool: - return _dtype_defs[t].numeric_category == NumericDTypeCategory.Float - - -def is_restricted_float(t: DType) -> bool: - return _dtype_defs[t].numeric_category == NumericDTypeCategory.RestrictedFloat - - -def is_arithmetic(t: DType) -> bool: - """Returns True if the |dtype| supports general arithmetic operations such as - addition, subtraction, multiplication, and division.""" - cat = _dtype_defs[t].numeric_category - return cat is not None and cat.arithmetic - - def broadcast_shapes(s1: Tuple[int, ...], s2: Tuple[int, ...]) -> Tuple[int, ...]: if len(s1) > len(s2): s1, s2 = s2, s1 @@ -638,31 +443,30 @@ class PointerInfo: @stub(host=True) def __init__(self, dtype: DType): - definition = _dtype_defs[dtype] - if not isinstance(definition, _PointerDTypeDefinition): + if not is_pointer_dtype(dtype): raise TypeError(f"'{dtype}' is not a pointer dtype") self._dtype = dtype - self._definition = definition @property @stub(host=True) def opaque(self) -> bool: """Whether the pointer dtype is opaque.""" - return self._definition.pointee_dtype is None + return _pointer_pointee_dtype(self._dtype) is None @property @stub(host=True) def pointee_dtype(self) -> DType: """Data type pointed to by this pointer dtype.""" - if self._definition.pointee_dtype is None: + ret = _pointer_pointee_dtype(self._dtype) + if ret is None: raise ValueError("Opaque pointer has no pointee dtype") - return self._definition.pointee_dtype + return ret @property @stub(host=True) def memory_space(self) -> MemorySpace: """CUDA memory space encoded in this pointer dtype.""" - return self._definition.memory_space + return _pointer_memory_space(self._dtype) def __repr__(self): if self.opaque: @@ -687,7 +491,7 @@ def __hash__(self): @stub(host=True) def is_pointer_dtype(dtype: DType) -> bool: """Return whether ``dtype`` is a pointer dtype.""" - return isinstance(_dtype_defs[dtype], _PointerDTypeDefinition) + return _is_pointer_dtype(dtype) @stub(host=True) @@ -716,46 +520,15 @@ def opaque_pointer_dtype(memory_space: MemorySpace = MemorySpace.GENERIC) -> DTy return _get_pointer_dtype(None, memory_space) -@cache -def _get_pointer_dtype(pointee_dtype: DType | None, memory_space: MemorySpace) -> DType: - match memory_space: - case MemorySpace.SHARED | MemorySpace.TENSOR | MemorySpace.SHARED_CLUSTER: - bitwidth = 32 - case _: - bitwidth = 64 - - params = [] - if pointee_dtype is None: - name = "opaque_pointer" - else: - assert isinstance(pointee_dtype, DType) - name = "pointer" - params.append(str(pointee_dtype)) - - if memory_space != MemorySpace.GENERIC: - params.append(f"MemorySpace.{memory_space._name_}") - - if len(params) > 0: - name += "[" + ", ".join(params) + "]" - - return _define_dtype(name, - _PointerDTypeDefinition(bitwidth=bitwidth, - pointee_dtype=pointee_dtype, - memory_space=memory_space)) - - # ============== Foreign Pointer DType =============== def is_foreign_pointer_dtype(dtype: DType) -> bool: - return isinstance(_dtype_defs[dtype], _ForeignPointerDTypeDefinition) + return _is_foreign_pointer_dtype(dtype) def foreign_pointer_pointee_dtype(dtype: DType) -> DType: """Return the pointee dtype encoded in a foreign pointer dtype.""" - definition = _dtype_defs[dtype] - if not isinstance(definition, _ForeignPointerDTypeDefinition): - raise TypeError(f"'{dtype}' is not a foreign pointer dtype") - return definition.pointee_dtype + return _foreign_pointer_pointee_dtype(dtype) @stub(host=True, static_eval_ok=True) @@ -764,11 +537,4 @@ def foreign_pointer_dtype(pointee_dtype: DType) -> DType: raise TypeError("pointee_dtype must be a cuda.tile dtype") if is_foreign_pointer_dtype(pointee_dtype): raise TypeError("nested foreign pointer dtypes are not supported") - - return _define_dtype( - f"foreign_pointer[{pointee_dtype}]", - _ForeignPointerDTypeDefinition( - bitwidth=64, - pointee_dtype=pointee_dtype, - ), - ) + return _get_foreign_pointer_dtype(pointee_dtype) diff --git a/test/test_ir_types.py b/test/test_ir_types.py index 87c07513..ac1a117f 100644 --- a/test/test_ir_types.py +++ b/test/test_ir_types.py @@ -21,7 +21,7 @@ uint64, uint32, uint16, uint8, bfloat16, tfloat32, float8_e4m3fn, float8_e5m2, is_boolean, is_integral, is_float, is_unrestricted_float, is_restricted_float, is_signed, - IntegerInfo, opaque_pointer_dtype, pointer_dtype, PointerInfo, + IntegerInfo, opaque_pointer_dtype, pointer_dtype, PointerInfo, foreign_pointer_dtype, ) from cuda.tile._ir.ops_utils import promote_dtypes from cuda.tile._ir.typing_support import to_dtype @@ -68,9 +68,20 @@ def test_builtin_types(): # Pickle-unpickle roundtrip assert pickle.loads(pickle.dumps(float16)) is float16 + assert pickle.loads(pickle.dumps(pointer_dtype(float16))) is pointer_dtype(float16) + assert (pickle.loads(pickle.dumps(pointer_dtype(float16, MemorySpace.SHARED))) + is pointer_dtype(float16, MemorySpace.SHARED)) + assert (pickle.loads(pickle.dumps(pointer_dtype(opaque_pointer_dtype(), MemorySpace.SHARED))) + is pointer_dtype(opaque_pointer_dtype(), MemorySpace.SHARED)) + assert (pickle.loads(pickle.dumps(opaque_pointer_dtype())) + is opaque_pointer_dtype()) + assert (pickle.loads(pickle.dumps(foreign_pointer_dtype(float16))) + is foreign_pointer_dtype(float16)) # Deep copy roundtrip assert copy.deepcopy(float16) is float16 + assert (copy.deepcopy(pointer_dtype(float16, MemorySpace.SHARED)) + is pointer_dtype(float16, MemorySpace.SHARED)) def test_tuple_type(): @@ -358,3 +369,7 @@ def test_pointer_info_equality(): for i, a in enumerate(dtypes): for j, b in enumerate(dtypes): assert (PointerInfo(a) == PointerInfo(b)) == (i == j) + + +def test_dtype_callable(): + assert callable(int8) From c1e9ab12c8930973091d4efe57be928a4c380e68 Mon Sep 17 00:00:00 2001 From: Jay Gu Date: Wed, 7 Oct 2026 09:02:22 -0700 Subject: [PATCH 3/7] [lang] Support bitcast in host compilation Also added dtype related function in compiled host. Signed-off-by: Jay Gu --- .../cuda-lang/src/cuda/lang/_stub/core_api.py | 4 ++-- experimental/cuda-lang/test/test_bitcast.py | 21 +++++++++++++++++++ experimental/cuda-lang/test/test_dtype.py | 12 +++++++++++ src/cuda/tile/_datatype.py | 6 +++--- 4 files changed, 38 insertions(+), 5 deletions(-) diff --git a/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py b/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py index 0d4e1c15..0f065d81 100644 --- a/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py +++ b/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py @@ -118,7 +118,7 @@ def __getitem__(self, indices: int | tuple[int, ...]) -> T: ... -@stub(host=True) +@stub(host=True, compiled_host=True) def dtype_of(value, /) -> DType: """ Returns the data type of a scalar, pointer, or vector value. @@ -681,7 +681,7 @@ def grid_dependency_control_launch_dependents() -> None: """Launch dependent grids in a programmatic dependent launch.""" -@stub +@stub(compiled_host=True) def bitcast(x, /, dtype): """Reinterpret a value as being of specified data type. """ diff --git a/experimental/cuda-lang/test/test_bitcast.py b/experimental/cuda-lang/test/test_bitcast.py index 8d3a7503..fe70ffff 100644 --- a/experimental/cuda-lang/test/test_bitcast.py +++ b/experimental/cuda-lang/test/test_bitcast.py @@ -8,6 +8,7 @@ import pytest import torch +from cuda.tile._cext import cconv_v3_enabled from cuda.lang._exception import CompilerExecutionError, TypeCheckingError from cuda.lang.compilation import KernelSignature from .util import filecheck, make_symbolic_tensor @@ -125,6 +126,26 @@ def kernel(out): assert got == 0xDEADBEEF, f"0x{got:x}" +@pytest.mark.skipif(not cconv_v3_enabled(), reason="Requires cconv3 enabled") +def test_bitcast_in_host_entry(): + @cl.kernel + def read_uint16(pointer, output): + cl.static_assert(pointer.pointee_dtype == cl.uint16) + output[0] = pointer[0] + + @cl.host_entry + def launcher(source, output): + pointer = cl.bitcast(source.pointer(), cl.pointer_dtype(cl.uint16)) + cl.launch(None, (1,), (1,), read_uint16, (pointer, output)) + + source = torch.tensor([0x34, 0x12], dtype=torch.uint8, device="cuda").view( + torch.float8_e4m3fn + ) + output = torch.zeros(1, dtype=torch.uint16, device="cuda") + launcher(source, output) + assert output.item() == 0x1234 + + @pytest.mark.parametrize("from_mspace", cl.MemorySpace._member_map_.values()) @pytest.mark.parametrize("to_mspace", cl.MemorySpace._member_map_.values()) def test_pointer_address_space_bitcast_compile_only(from_mspace, to_mspace): diff --git a/experimental/cuda-lang/test/test_dtype.py b/experimental/cuda-lang/test/test_dtype.py index 414bb7d7..5a6cbddb 100644 --- a/experimental/cuda-lang/test/test_dtype.py +++ b/experimental/cuda-lang/test/test_dtype.py @@ -80,6 +80,18 @@ def kern(x: cl.Array): cl.launch(torch.cuda.current_stream(), (1,), (1,), kern, (x,)) +def test_dtype_helpers_in_host_entry(): + @cl.host_entry + def inspect(value): + cl.static_assert(cl.dtype_of(value) == cl.int32) + dtype = cl.pointer_dtype(cl.uint16) + cl.static_assert(cl.is_pointer_dtype(dtype)) + cl.static_assert(cl.is_pointer_dtype(cl.opaque_pointer_dtype())) + cl.static_assert(cl.PointerInfo(dtype).pointee_dtype == cl.uint16) + + inspect(42) + + @pytest.mark.parametrize( "memory_space, bitwidth", ( diff --git a/src/cuda/tile/_datatype.py b/src/cuda/tile/_datatype.py index 5f3d4104..ad6865dc 100644 --- a/src/cuda/tile/_datatype.py +++ b/src/cuda/tile/_datatype.py @@ -488,13 +488,13 @@ def __hash__(self): return hash(self._dtype) -@stub(host=True) +@stub(host=True, compiled_host=True) def is_pointer_dtype(dtype: DType) -> bool: """Return whether ``dtype`` is a pointer dtype.""" return _is_pointer_dtype(dtype) -@stub(host=True) +@stub(host=True, compiled_host=True) def pointer_dtype(pointee_dtype: DType, memory_space: MemorySpace = MemorySpace.GENERIC) -> DType: """Return the dtype for a pointer to ``pointee_dtype`` in ``memory_space``. @@ -509,7 +509,7 @@ def pointer_dtype(pointee_dtype: DType, return _get_pointer_dtype(pointee_dtype, memory_space) -@stub(host=True) +@stub(host=True, compiled_host=True) def opaque_pointer_dtype(memory_space: MemorySpace = MemorySpace.GENERIC) -> DType: """Return the dtype for an opaque pointer in ``memory_space``. From d6fea7a288f52d58d71670d5ed8d2dbfb246300c Mon Sep 17 00:00:00 2001 From: Greg Bonik Date: Wed, 30 Sep 2026 17:25:38 -0700 Subject: [PATCH 4/7] Replace g_spy_mutex with GlobalLock Signed-off-by: Greg Bonik --- cext/cuda_helper.cpp | 30 ++++++++++++------------------ 1 file changed, 12 insertions(+), 18 deletions(-) diff --git a/cext/cuda_helper.cpp b/cext/cuda_helper.cpp index f197b80c..d84def42 100644 --- a/cext/cuda_helper.cpp +++ b/cext/cuda_helper.cpp @@ -202,23 +202,17 @@ PyObject* destroy_stream(PyObject* self, PyObject* arg) { Py_RETURN_NONE; } -static decltype(cuLaunchKernelEx)* g_real_cuLaunchKernelEx; // Protected by the GIL or g_spy_mutex -static PyObject* g_cuLaunchKernelEx_spy_callback; // Protected by the GIL or g_spy_mutex - -#ifdef Py_GIL_DISABLED -static PyMutex g_spy_mutex = {0}; -#endif +static ProtectedByGlobalLock g_real_cuLaunchKernelEx; +static ProtectedByGlobalLock g_cuLaunchKernelEx_spy_callback; static CUresult shim_cuLaunchKernelEx( const CUlaunchConfig *config, CUfunction f, void** kernelParams, void** extra) { -#ifdef Py_GIL_DISABLED - PyCriticalSectionGuard guard(&g_spy_mutex); -#endif + GlobalLock lock; PyPtr res = steal(PyObject_CallFunction( - g_cuLaunchKernelEx_spy_callback, + g_cuLaunchKernelEx_spy_callback.get(lock), "(K III III I K)", reinterpret_cast(f), config->gridDimX, config->gridDimY, config->gridDimZ, @@ -228,13 +222,13 @@ static CUresult shim_cuLaunchKernelEx( )); if (!res) return CUDA_ERROR_LAUNCH_FAILED; - return g_real_cuLaunchKernelEx(config, f, kernelParams, extra); + return g_real_cuLaunchKernelEx.get(lock)(config, f, kernelParams, extra); } static PyObject* spy_on_cuLaunchKernel_begin(PyObject* self, PyObject* arg) { GlobalLock lock; - if (g_real_cuLaunchKernelEx) { + if (g_real_cuLaunchKernelEx.get(lock)) { raise(PyExc_RuntimeError, "Already spying"); return nullptr; } @@ -243,8 +237,8 @@ static PyObject* spy_on_cuLaunchKernel_begin(PyObject* self, PyObject* arg) { if (!driver_result.is_ok()) return nullptr; DriverApi* api = const_cast(*driver_result); - g_real_cuLaunchKernelEx = api->cuLaunchKernelEx; - g_cuLaunchKernelEx_spy_callback = Py_NewRef(arg); + g_real_cuLaunchKernelEx.get(lock) = api->cuLaunchKernelEx; + g_cuLaunchKernelEx_spy_callback.get(lock) = Py_NewRef(arg); api->cuLaunchKernelEx = shim_cuLaunchKernelEx; return Py_NewRef(Py_None); } @@ -252,7 +246,7 @@ static PyObject* spy_on_cuLaunchKernel_begin(PyObject* self, PyObject* arg) { static PyObject* spy_on_cuLaunchKernel_end(PyObject* self, PyObject* arg) { GlobalLock lock; - if (!g_real_cuLaunchKernelEx) { + if (!g_real_cuLaunchKernelEx.get(lock)) { raise(PyExc_RuntimeError, "Not spying"); return nullptr; } @@ -261,9 +255,9 @@ static PyObject* spy_on_cuLaunchKernel_end(PyObject* self, PyObject* arg) { if (!driver_result.is_ok()) return nullptr; DriverApi* api = const_cast(*driver_result); - api->cuLaunchKernelEx = g_real_cuLaunchKernelEx; - g_real_cuLaunchKernelEx = nullptr; - Py_CLEAR(g_cuLaunchKernelEx_spy_callback); + api->cuLaunchKernelEx = g_real_cuLaunchKernelEx.get(lock); + g_real_cuLaunchKernelEx.get(lock) = nullptr; + Py_CLEAR(g_cuLaunchKernelEx_spy_callback.get(lock)); return Py_NewRef(Py_None); } From 9550ed6b8bf927d6d6c379327aa40150b7e5dd09 Mon Sep 17 00:00:00 2001 From: Gideon Kassa Date: Wed, 7 Oct 2026 17:06:52 -0700 Subject: [PATCH 5/7] * [lang] Add support for array slicing Signed-off-by: Gideon Kassa --- .../src/cuda/lang/_ir/op_impl/pointer_impl.py | 151 ++++++++++++++-- .../src/cuda/lang/_ir/op_impl/vector_impl.py | 5 +- .../cuda-lang/src/cuda/lang/_ir/ops.py | 13 ++ .../cuda-lang/src/cuda/lang/_ir/type.py | 4 + .../cuda-lang/src/cuda/lang/_stub/core_api.py | 32 +++- experimental/cuda-lang/test/test_slice.py | 170 ++++++++++++++++++ experimental/cuda-lang/test/test_vectors.py | 9 +- src/cuda/tile/_ir/core_ops.py | 14 +- src/cuda/tile/_ir/ops.py | 6 +- src/cuda/tile/_ir/type.py | 35 +++- 10 files changed, 401 insertions(+), 38 deletions(-) create mode 100644 experimental/cuda-lang/test/test_slice.py diff --git a/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/pointer_impl.py b/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/pointer_impl.py index 7a78fedb..9fb8b6d9 100644 --- a/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/pointer_impl.py +++ b/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/pointer_impl.py @@ -18,6 +18,9 @@ Type, VectorTy, make_rank0_ty, + SliceType, + TupleTy, + EllipsisType, ) from cuda.lang._ir.atomics_support import ( require_atomic_memory_order_and_scope, @@ -44,9 +47,13 @@ pointer_dtype, uint64, ) -from cuda.tile._ir.arithmetic_ops import astype, binary_arithmetic_tensorlike_raw +from cuda.tile._ir.arithmetic_ops import ( + astype, binary_arithmetic_tensorlike_raw, binary_arithmetic_tensorlike +) from cuda.tile._ir.cast_ops import address_space_cast, implicit_cast -from cuda.tile._ir.core_ops import bind_method, loosely_typed_const, strictly_typed_const +from cuda.tile._ir.core_ops import ( + bind_method, loosely_typed_const, strictly_typed_const, build_slice, build_tuple, +) from cuda.tile._ir.ir import add_operation_variadic, make_aggregate from cuda.tile._ir.op_impl import ( ImplRegistry, @@ -257,19 +264,6 @@ def pointer_setitem(object: Var[PointerTy], key: Var[Type], value: Var[Type]): ) -@impl(operator.getitem, overload=(ArrayTy, WILDCARD)) -def array_getitem(object: Var, key: Var) -> Var: - array_ty = require_array_type(object) - indices = require_array_indices(object, key) - pointer = _array_element_pointer(object, indices) - return add_operation( - LoadPointer, - make_rank0_ty(array_ty.dtype), - pointer=pointer, - alignment=None, - ) - - @impl(operator.setitem, overload=(ArrayTy, WILDCARD, WILDCARD)) def array_setitem(object: Var, key: Var, value: Var): array_ty = require_array_type(object) @@ -292,6 +286,133 @@ def array_setitem(object: Var, key: Var, value: Var): ) +def _require_nonnegative_constant(var: Var, message: str) -> None: + if var.is_constant() and var.get_constant() < 0: + raise TypeCheckingError(message) + return var + + +def _require_valid_slice_constant(start: Var, stop: Var, step: Var, array_dim: Var) -> None: + start = _require_nonnegative_constant(start, "A non-negative slice start is required") + stop = _require_nonnegative_constant(stop, "A non-negative slice stop is required") + if step.is_constant() and step.get_constant() < 1: + raise TypeCheckingError("A positive slice step is required") + if stop.is_constant(): + stop_c = stop.get_constant() + if start.is_constant(): + if start.get_constant() >= stop_c: + raise TypeCheckingError("View cannot be empty") + if array_dim.is_constant(): + array_dim_c = array_dim.get_constant() + if stop_c > array_dim_c: + raise TypeCheckingError(f"The provided slice stop ({stop_c}) is greater " + f"than the array dimension ({array_dim_c})") + return start, stop, step, array_dim + + +@impl(operator.getitem, overload=(ArrayTy, WILDCARD)) +def array_getitem(object: Var, key: Var) -> Var: + original_array_ty = require_array_type(object) + original_array_val = object.get_aggregate() + + original_array_shape = original_array_val.shape + original_array_strides = original_array_val.strides + + new_shape = tuple() + new_strides = tuple() + offset = strictly_typed_const(0, ScalarTy(uint64)) + + key_ty = key.get_type() + if isinstance(key_ty, TupleTy): + tuple_value = key.get_aggregate() + indices = tuple_value.items + else: + # Normalize everything to tuples + indices = (key,) + + num_slices = sum(isinstance(i.get_type(), SliceType) for i in indices) + num_ellipsis = sum(isinstance(i.get_type(), EllipsisType) for i in indices) + + if (num_ellipsis + num_slices == 0) and len(indices) == original_array_ty.ndim: + indices = tuple(implicit_cast(var, original_array_ty.index_dtype, + "Invalid array index") for var in indices) + pointer = _array_element_pointer(object, indices) + return add_operation( + LoadPointer, + make_rank0_ty(original_array_ty.dtype), + pointer=pointer, + alignment=None, + ) + + if num_ellipsis > 1: + raise TypeCheckingError("An index can only have a single ellipsis ('...')") + + if (len(indices) - num_ellipsis > original_array_ty.ndim): + raise TypeCheckingError( + "There were more indices used in slice than there are dimensions in the array" + ) + + new_indices = [] + for index in indices: + if isinstance(index.get_type(), EllipsisType): + padding_size = original_array_ty.ndim - len(indices) + 1 + new_indices += [None] * padding_size + elif isinstance(index.get_type(), SliceType): + array_dim = original_array_shape[len(new_indices)] + s = index.get_aggregate() + start, stop, step = s.start, s.stop, s.step + + start = loosely_typed_const(0) if is_none(start) else start + stop = array_dim if is_none(stop) else stop + step = loosely_typed_const(1) if is_none(step) else step + start, stop, step, array_dim = _require_valid_slice_constant( + start, stop, step, array_dim + ) + start, stop, step = ( + implicit_cast(j, original_array_ty.index_dtype, "Invalid slice index dtype") + for j in (start, stop, step) + ) + new_indices.append(build_slice((start, stop, step))) + elif isinstance(index.get_type(), ScalarTy): + new_indices.append(index) + else: + raise TypeCheckingError("The provided slice index is not valid") + + indices = new_indices + [None for _ in range(original_array_ty.ndim - len(new_indices))] + for index, stride, array_dim in zip(indices, original_array_strides, + original_array_shape, strict=True): + if index is None: + new_shape += (array_dim,) + new_strides += (stride,) + index = loosely_typed_const(0) + elif isinstance(index.get_type(), SliceType): + s = index.get_aggregate() + start, stop, step = s.start, s.stop, s.step + delta = binary_arithmetic_tensorlike("sub", stop, start) + new_shape += (binary_arithmetic_tensorlike("cdiv", delta, step),) + new_strides += (binary_arithmetic_tensorlike("mul", stride, step),) + index = start + elif isinstance(index.get_type(), ScalarTy): + index = implicit_cast(index, original_array_ty.index_dtype, "Invalid array index") + + scaled = binary_arithmetic_tensorlike( + "mul", + astype(index, datatype.uint64), + astype(stride, datatype.uint64) + ) + + offset = binary_arithmetic_tensorlike("add", offset, scaled) + + original_array_base_ptr = _get_array_base_pointer(object) + pointer = add_operation( + PointerOffset, + original_array_base_ptr.get_type(), + pointer=original_array_base_ptr, + offset=offset, + ) + return array_from_parts_impl(pointer, build_tuple(new_shape), build_tuple(new_strides)) + + def load_pointer(pointer: Var[PointerTy], count: int | None = None, alignment: int | None = None): pointee_dtype = pointer.get_type().pointee_dtype if count is None or count == 1: diff --git a/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/vector_impl.py b/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/vector_impl.py index 7be70dc4..d122ebf0 100644 --- a/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/vector_impl.py +++ b/experimental/cuda-lang/src/cuda/lang/_ir/op_impl/vector_impl.py @@ -16,7 +16,7 @@ from cuda.tile._ir.arithmetic_ops import astype from cuda.tile._ir.cast_ops import implicit_cast from cuda.tile._ir.core_ops import bind_method, build_tuple, loosely_typed_const -from cuda.tile._ir.ops import strictly_typed_const, slice_impl +from cuda.tile._ir.ops import strictly_typed_const from cuda.tile._ir.ops_utils import promote_dtypes from cuda.tile._ir.type import LooselyTypedScalar from cuda.lang._exception import InternalError, TypeCheckingError, InvalidValueError @@ -122,9 +122,6 @@ def vector_constructor_impl(elements: tuple[Var, ...], dtype: Var) -> Var[Vector return vector_construct(VectorTy(element_dtype, len(values)), values) -impl(slice)(slice_impl) - - @impl(operator.getitem, overload=(VectorTy, SliceType)) def vector_slice_impl(object: Var[VectorTy], key: Var[SliceType]): s = require_constant_slice(key) diff --git a/experimental/cuda-lang/src/cuda/lang/_ir/ops.py b/experimental/cuda-lang/src/cuda/lang/_ir/ops.py index 29ccc1ed..6b5462f9 100644 --- a/experimental/cuda-lang/src/cuda/lang/_ir/ops.py +++ b/experimental/cuda-lang/src/cuda/lang/_ir/ops.py @@ -131,6 +131,8 @@ PointerTy, VectorTy, DTypeSpec, + SliceValue, + SliceType, ) from .ir import ( @@ -405,6 +407,17 @@ def vector_dtype_impl(object: Var[VectorTy], name: Var): return loosely_typed_const(object.get_type().element_dtype) +@impl(slice) +def slice_impl(start: Var, stop: Var, step: Var) -> Var: + res = make_aggregate( + SliceValue(start, stop, step), + SliceType((start.get_type(), stop.get_type(), step.get_type())) + ) + if (start.is_constant() and stop.is_constant() and step.is_constant()): + res.set_constant(slice(start.get_constant(), stop.get_constant(), step.get_constant())) + return res + + @dataclass(eq=False) class Branch(Operation, opcode="br", terminator=True): target: Block = successor() diff --git a/experimental/cuda-lang/src/cuda/lang/_ir/type.py b/experimental/cuda-lang/src/cuda/lang/_ir/type.py index b2d9d2d9..2d7050ca 100644 --- a/experimental/cuda-lang/src/cuda/lang/_ir/type.py +++ b/experimental/cuda-lang/src/cuda/lang/_ir/type.py @@ -33,6 +33,8 @@ SymbolicClosure, SliceType, StreamTy, + SliceValue, + EllipsisType, ) import cuda.lang._datatype as datatype from cuda.tile._datatype import DType, PointerInfo @@ -322,4 +324,6 @@ def type_bitwidth(x: Type): "SymbolicScalar", "SymbolicPointer", "SliceType", + "SliceValue", + "EllipsisType", ) diff --git a/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py b/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py index 0f065d81..06c13349 100644 --- a/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py +++ b/experimental/cuda-lang/src/cuda/lang/_stub/core_api.py @@ -32,6 +32,30 @@ def __exit__(self, exc_type, exc_val, exc_tb): ... class Array(TileArray, Generic[T]): """ N-dimensional array type. + + Indexing every dimension with an integer accesses one element. + Partial indexing or slices return an ``Array`` view. + + Examples: + + .. testcode:: + :template: kernel_wrapper.py + + array = cl.shared_array(4, cl.int32) + + # Integer indexing. + array[1] = 7 + print(array[1]) + + # Slice indexing returns a view. + array_view = array[1:3] + array_view[0] = 9 + print(array[1]) + + .. testoutput:: + + 7 + 9 """ @staticmethod @@ -111,9 +135,11 @@ def __setitem__(self, indices: int | tuple[int, ...], value: T): ... @stub - def __getitem__(self, indices: int | tuple[int, ...]) -> T: - """Retrieve the value given by ``indices``. - Equivalent to ``self.pointer(indices).load()``. + def __getitem__(self, indices: int | slice | tuple[int | slice, ...]) -> T | "Array[T]": + """Retrieve value or view given by ``indices``. + When every dimension is indexed by an integer, this is equivalent to + ``self.pointer(indices).load()``. + Partial indexing or slices return an ``Array`` view. """ ... diff --git a/experimental/cuda-lang/test/test_slice.py b/experimental/cuda-lang/test/test_slice.py new file mode 100644 index 00000000..1ca658ac --- /dev/null +++ b/experimental/cuda-lang/test/test_slice.py @@ -0,0 +1,170 @@ +# SPDX-FileCopyrightText: Copyright (c) <2026> NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# SPDX-License-Identifier: Apache-2.0 + +import cuda.lang as cl +import torch +import pytest +from cuda.lang._exception import TypeCheckingError + + +@pytest.mark.parametrize("index, result", + [((None, 4, None), [0, 1, 2, 3]), + ((4, 8, None), [4, 5, 6, 7]), + ((4, 12, 2), [4, 6, 8, 10])]) +def test_array_slice_view(index, result): + @cl.kernel + def kernel(a, b): + i = slice(index[0], index[1], index[2]) + view = a[i] + b[0] = view[0] + b[1] = view[1] + b[2] = view[2] + b[3] = view[3] + + a = torch.arange(16, dtype=torch.int32).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b)) + assert b.cpu().tolist() == result + + +def test_array_slice_view_2d_array(): + @cl.kernel + def kernel(a, b): + view = a[1, 4:] + b[0] = view[0] + b[1] = view[1] + b[2] = view[2] + b[3] = view[3] + + a = torch.arange(16, dtype=torch.int32).reshape((2, 8)).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b)) + assert b.cpu().tolist() == [12, 13, 14, 15] + + +def test_array_slice_view_runtime_slice(): + @cl.kernel + def kernel(a, b, i): + start, step = i[0], i[1] + view = a[start::step] + b[0] = view[0] + b[1] = view[1] + b[2] = view[2] + b[3] = view[3] + + a = torch.arange(16, dtype=torch.int32).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + i = torch.tensor([8, 2], dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b, i)) + assert b.cpu().tolist() == [8, 10, 12, 14] + + +def test_array_slice_view_static_shape(): + @cl.kernel + def kernel(a): + b = a[:, :] + view = b[:3:1, :] + cl.static_assert(view.shape[0] == 3) + + a = torch.arange(16, dtype=torch.int32).reshape((4, 4)).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a,)) + + +def test_array_slice_view_stored_slice(): + @cl.kernel + def kernel(a, b, i): + if i[0] > 1: + s = slice(i[1], None, 1) + else: + s = slice(i[0], None, 1) + view = a[s] + b[0] = view[0] + b[1] = view[1] + b[2] = view[2] + b[3] = view[3] + + a = torch.arange(16, dtype=torch.int32).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + i = torch.tensor([8, 2], dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b, i)) + assert b.cpu().tolist() == [2, 3, 4, 5] + + +def test_array_slice_view_external_slice(): + s = slice(1, 3, 1) + + @cl.kernel + def kernel(a, b): + view = a[s, ...] + b[0] = view[0, 0] + b[1] = view[1, 0] + b[2] = view[1, 1] + b[3] = view[0, 1] + + a = torch.arange(16, dtype=torch.int32).reshape((4, 4)).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b)) + assert b.cpu().tolist() == [4, 8, 9, 5] + + +def test_array_slice_view_ellipsis(): + + @cl.kernel + def kernel(a, b): + view = a[0, ..., 3] + view = view[..., :, :] + b[0] = view[0, 0] + b[1] = view[1, 1] + b[2] = view[2, 2] + b[3] = view[3, 3] + + a = torch.arange(256, dtype=torch.int32).reshape((4, 4, 4, 4)).cuda() + b = torch.zeros(4, dtype=torch.int32).cuda() + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a, b)) + assert b.cpu().tolist() == [3, 23, 43, 63] + + +@pytest.mark.parametrize("start, stop, step, message", + [(-2, None, None, "A non-negative slice start is required"), + (None, -1, None, "A non-negative slice stop is required"), + (None, None, 0, "A positive slice step is required")]) +def test_array_slice_view_static_index_reject(start, stop, step, message): + @cl.kernel + def kernel(a): + a[start:stop:step] + + a = torch.arange(16, dtype=torch.int32).cuda() + with pytest.raises(TypeCheckingError, match=message): + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a,)) + + +def test_array_slice_view_static_array_bounds_reject(): + @cl.kernel + def kernel(): + a = cl.shared_array(16, cl.int32) + a[1:18:2] + + with pytest.raises(TypeCheckingError, match="The provided slice stop"): + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, ()) + + +def test_array_slice_view_static_view_bounds_reject(): + @cl.kernel + def kernel(a): + view = a[12:14] + view[:3] + + a = torch.arange(16, dtype=torch.int32).cuda() + with pytest.raises(TypeCheckingError, match="The provided slice stop"): + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a,)) + + +def test_array_slice_view_none_expand_reject(): + @cl.kernel + def kernel(a): + a[None, :] + + a = torch.arange(16, dtype=torch.int32).reshape(4, 4).cuda() + with pytest.raises(TypeCheckingError, match="The provided slice index is not valid"): + cl.launch(torch.cuda.current_stream(), (1,), (1,), kernel, (a,)) diff --git a/experimental/cuda-lang/test/test_vectors.py b/experimental/cuda-lang/test/test_vectors.py index 82e0525f..9eb0fc95 100644 --- a/experimental/cuda-lang/test/test_vectors.py +++ b/experimental/cuda-lang/test/test_vectors.py @@ -923,7 +923,8 @@ def kernel(): compile_kernel( kernel, raises=pytest.raises( - TypeCheckingError, match="Non-constant slices are not supported" + TypeCheckingError, + match="Expected a slice constant, but given value is not constant" ), ) @@ -935,7 +936,8 @@ def kernel(): compile_kernel( kernel, raises=pytest.raises( - TypeCheckingError, match="Non-constant slices are not supported" + TypeCheckingError, + match="Expected a slice constant, but given value is not constant" ), ) @@ -947,7 +949,8 @@ def kernel(): compile_kernel( kernel, raises=pytest.raises( - TypeCheckingError, match="Non-constant slices are not supported" + TypeCheckingError, + match="Expected a slice constant, but given value is not constant" ), ) diff --git a/src/cuda/tile/_ir/core_ops.py b/src/cuda/tile/_ir/core_ops.py index d8e0bf8e..54974d40 100644 --- a/src/cuda/tile/_ir/core_ops.py +++ b/src/cuda/tile/_ir/core_ops.py @@ -35,7 +35,7 @@ RangeIterType, RangeValue, TypeTy, ModuleTy, NONE, SliceType, StringTy, FormattedStringTy, \ StringFormat, FormattedStringValue, FormattedPiece, DictTy, DictValue, EnumTy, TokenTy, \ FunctionTy, GeneratorContextManagerTy, ContextManagerState, GeneratorContextManagerValue, \ - NotImplementedTy + NotImplementedTy, SliceValue from cuda.tile._ir.typing_support import type_of_constant_python_value, \ loose_type_of_constant_python_value, get_dataclass_info, \ create_dataclass_instance, find_method, dataclass_has_default_repr, dataclass_has_default_cmp @@ -905,6 +905,14 @@ async def len_dataclass_impl(x: Var[DataclassTy]) -> Var: # =========================================================================================== +def build_slice(items: Sequence[Var]) -> Var: + ty = SliceType(tuple(x.get_type() for x in items)) + res = make_aggregate(SliceValue(*items), ty) + if all(i.is_constant() for i in items): + res.set_constant(slice(*(i.get_constant() for i in items))) + return res + + def bind_method(object: Var, func) -> Var: agg_value = BoundMethodValue(object) res_ty = BoundMethodTy(object.get_type(), func) @@ -922,6 +930,10 @@ def sym2var(x: Any, constant_only: bool = False) -> Var: if isinstance(x, tuple): return build_tuple(tuple(sym2var(item, constant_only=constant_only) for item in x)) + if isinstance(x, slice): + return build_slice(tuple(sym2var(i, constant_only=constant_only) + for i in (x.start, x.stop, x.step))) + cls = type(x) if dataclasses.is_dataclass(cls): info = get_dataclass_info(cls) diff --git a/src/cuda/tile/_ir/ops.py b/src/cuda/tile/_ir/ops.py index 89e3824a..24cff7ab 100644 --- a/src/cuda/tile/_ir/ops.py +++ b/src/cuda/tile/_ir/ops.py @@ -32,7 +32,8 @@ from .cast_ops import implicit_cast from .control_flow_ops import Loop, IfElse, control_flow_impl_registry, EndBranch from .core_ops import loosely_typed_const, strictly_typed_const, build_tuple, bind_method, \ - sym2var, core_impl_registry, print_impl, TilePrintf, tuple_item, comparison_operator_impl + sym2var, core_impl_registry, print_impl, TilePrintf, tuple_item, comparison_operator_impl, \ + build_slice from .static_eval_ops import static_eval_impl_registry from .type import ( TupleValue, ArrayValue, ListValue, TiledViewValue, RawArrayMemoryValue, @@ -225,8 +226,7 @@ def tile_divmod_function_impl(x: Var, y: Var): def slice_impl(start: Var, stop: Var, step: Var) -> Var: if not (start.is_constant() and stop.is_constant() and step.is_constant()): raise TileTypeError("Non-constant slices are not supported") - return loosely_typed_const( - slice(start.get_constant(), stop.get_constant(), step.get_constant())) + return build_slice((start, stop, step)) # =========================================================================================== diff --git a/src/cuda/tile/_ir/type.py b/src/cuda/tile/_ir/type.py index 33c22c49..0823f49d 100644 --- a/src/cuda/tile/_ir/type.py +++ b/src/cuda/tile/_ir/type.py @@ -140,22 +140,39 @@ def __str__(self): # ============== Slice Type =============== +@dataclass(frozen=True) class SliceType(Type): - _instance = None + item_types: tuple["Type", "Type", "Type"] | None = None - def __new__(cls): - if cls._instance is None: - cls._instance = super().__new__(cls) - return cls._instance + def is_aggregate(self) -> bool: + return self.item_types is not None - def __str__(self): - return "Slice" + def aggregate_item_types(self) -> tuple["Type", "Type", "Type"] | None: + return self.item_types + + def make_aggregate_value(self, items: tuple["Type", "Type", "Type"]) -> "AggregateValue": + return SliceValue(*items) def __eq__(self, other: Type): - return isinstance(other, SliceType) + return isinstance(other, SliceType) and self.item_types == other.item_types def __hash__(self): - return hash("SliceType") + return hash(("SliceType", self.item_types)) + + def __str__(self): + if self.item_types is not None: + return 'Slice(' + ','.join(str(x) for x in self.item_types) + ')' + return 'Slice()' + + +@dataclass +class SliceValue(AggregateValue): + start: "Var" + stop: "Var" + step: "Var" + + def as_tuple(self) -> tuple["Var", ...]: + return (self.start, self.stop, self.step) SLICE = SliceType() From a7c807551833c95e04ede909907cee824ef34107 Mon Sep 17 00:00:00 2001 From: Ziheng Deng Date: Thu, 1 Oct 2026 11:54:13 -0700 Subject: [PATCH 6/7] Allow oversized shared memory for driver version R615+ Signed-off-by: Ziheng Deng --- cext/cuda_helper.cpp | 25 +++++--------------- cext/cuda_loader.cpp | 17 +++++++++++++ cext/cuda_loader.h | 5 ++++ cext/tile_kernel.cpp | 23 ++++++++++++------ changelog.d/allow-oversized-shared-memory.md | 2 ++ 5 files changed, 46 insertions(+), 26 deletions(-) create mode 100644 changelog.d/allow-oversized-shared-memory.md diff --git a/cext/cuda_helper.cpp b/cext/cuda_helper.cpp index d84def42..725aa617 100644 --- a/cext/cuda_helper.cpp +++ b/cext/cuda_helper.cpp @@ -16,14 +16,9 @@ const char* get_cuda_error(const DriverApi* driver, CUresult res) { } Status check_driver_version(const DriverApi* driver, int minimum_version) { - int version; - CUresult res = driver->cuDriverGetVersion(&version); - if (res != CUDA_SUCCESS) { - return raise(PyExc_RuntimeError, "cuDriverGetVersion: ", get_cuda_error(driver, res)); - } - if (version < minimum_version) { - int major = version / 1000; - int minor = (version % 1000) / 10; + if (driver->cuda_version < minimum_version) { + int major = driver->cuda_version / 1000; + int minor = (driver->cuda_version % 1000) / 10; int required_major = minimum_version / 1000; return raise(PyExc_RuntimeError, "Minimum driver version required is ", required_major, ".0, got ", @@ -135,20 +130,12 @@ PyObject* get_compute_capability(PyObject *self, PyObject *args) { } PyObject* get_driver_version(PyObject *self, PyObject *Py_UNUSED(ignored)) { - int major, minor; - GlobalLock lock; Result driver_result = get_driver_api(lock); if (!driver_result.is_ok()) return nullptr; - const DriverApi* d = *driver_result; - - CUresult res = d->cuDriverGetVersion(&major); - if (res != CUDA_SUCCESS) { - raise(PyExc_RuntimeError, "cuDriverGetVersion: ", get_cuda_error(d, res)); - return nullptr; - } - minor = (major % 1000) / 10; - major = major / 1000; + int version = (*driver_result)->cuda_version; + int major = version / 1000; + int minor = (version % 1000) / 10; return Py_BuildValue("(ii)", major, minor); } diff --git a/cext/cuda_loader.cpp b/cext/cuda_loader.cpp index e5a79fc7..61e9ccbe 100644 --- a/cext/cuda_loader.cpp +++ b/cext/cuda_loader.cpp @@ -52,9 +52,22 @@ FOREACH_CUDA_FUNCTION_TO_LOAD(DEFINE_CUDA_FUNCTION_GLOBAL) Status driver_api_init(DriverApi* driver_api, cuGetProcAddress_v2_t _cuGetProcAddress) { FOREACH_CUDA_FUNCTION_TO_LOAD(GET_PROC_ADDRESS) + driver_api->cuda_version = 0; return OK; } +std::optional DriverApi::get_shared_memory_mode_attribute() const { +#if CUDA_VERSION >= 13040 + if (cuda_version >= 13040) { + CUlaunchAttribute attribute = {}; + attribute.id = CU_LAUNCH_ATTRIBUTE_SHARED_MEMORY_MODE; + attribute.value.sharedMemoryMode = CU_SHARED_MEMORY_MODE_ALLOW_OVERSIZED_SHARED_MEMORY; + return attribute; + } +#endif + return std::nullopt; +} + static Result get_cuGetProcAddress_from_python() { PyPtr load_libcuda_mod = steal(PyImport_ImportModule("cuda.tile._load_libcuda")); if (!load_libcuda_mod) return ErrorRaised; @@ -85,6 +98,10 @@ Result get_driver_api(GlobalLock& lock) { CUresult res = instance.cuInit(0); if (res != CUDA_SUCCESS) return raise(PyExc_RuntimeError, "cuInit: ", get_cuda_error(&instance, res)); + res = instance.cuDriverGetVersion(&instance.cuda_version); + if (res != CUDA_SUCCESS) + return raise(PyExc_RuntimeError, "cuDriverGetVersion: ", + get_cuda_error(&instance, res)); if (!check_driver_version(&instance, MIN_DRIVER_VERSION)) return ErrorRaised; initialized = true; diff --git a/cext/cuda_loader.h b/cext/cuda_loader.h index 4487ca54..8f3ba28d 100644 --- a/cext/cuda_loader.h +++ b/cext/cuda_loader.h @@ -8,6 +8,7 @@ #include "py.h" #include +#include #define FOREACH_CUDA_FUNCTION_TO_LOAD(X) \ X(cuInit, "cuInit", 2000) \ @@ -59,6 +60,7 @@ X(cuGraphDestroy, "cuGraphDestroy", 10000) \ X(cuGraphAddEventRecordNode, "cuGraphAddEventRecordNode", 11010) \ X(cuGraphAddKernelNode, "cuGraphAddKernelNode", 12000) \ + X(cuGraphKernelNodeSetAttribute, "cuGraphKernelNodeSetAttribute", 11000) \ X(cuGraphAddMemsetNode, "cuGraphAddMemsetNode", 10000) \ X(cuGraphAddMemAllocNode, "cuGraphAddMemAllocNode", 11040) \ X(cuGraphAddMemFreeNode, "cuGraphAddMemFreeNode", 11040) \ @@ -73,6 +75,9 @@ struct DriverApi { FOREACH_CUDA_FUNCTION_TO_LOAD(DECLARE_CUDA_FUNC_EXTERN) + int cuda_version; + + std::optional get_shared_memory_mode_attribute() const; }; diff --git a/cext/tile_kernel.cpp b/cext/tile_kernel.cpp index 9ccbcbb7..49c7d57d 100644 --- a/cext/tile_kernel.cpp +++ b/cext/tile_kernel.cpp @@ -4822,6 +4822,11 @@ static Result benchmark(const DriverApi* driver, CUgraphNode kernel_node; CU_CHECK("cuGraphAddKernelNode", driver->cuGraphAddKernelNode(&kernel_node, graph.get(), &start_node, 1, &kparams)); + if (auto attribute = driver->get_shared_memory_mode_attribute()) { + CU_CHECK("cuGraphKernelNodeSetAttribute", + driver->cuGraphKernelNodeSetAttribute( + kernel_node, attribute->id, &attribute->value)); + } // Event: end of kernel CUgraphNode end_node; @@ -5293,18 +5298,21 @@ struct LaunchArgs { }; -static Result parse_tile_launch_kwargs(PyObject *const *args, +static Result parse_tile_launch_kwargs(const DriverApi* driver, + PyObject *const *args, Py_ssize_t nargs, PyObject *kwargs, CUlaunchAttribute launch_attrs[kMaxCUlaunchAttrs] ) { + size_t num_attrs = 0; + if (auto attribute = driver->get_shared_memory_mode_attribute()) + launch_attrs[num_attrs++] = *attribute; if (kwargs == nullptr) - return 0; + return num_attrs; CHECK(PyTuple_Check(kwargs) && "Keyword argument tuple is nonnull and not a tuple"); const auto nkwargs = PyTuple_GET_SIZE(kwargs); - size_t num_attrs = 0; for (Py_ssize_t i = 0; i < nkwargs; i++) { PyObject *keyword = PyTuple_GET_ITEM(kwargs, i); @@ -5315,6 +5323,7 @@ static Result parse_tile_launch_kwargs(PyObject *const *args, if (!PyBool_Check(kwarg)) return raise(PyExc_TypeError, "expected argument ", keyword, " to have type bool"); + CHECK(num_attrs < kMaxCUlaunchAttrs); CUlaunchAttribute *attr = &launch_attrs[num_attrs++]; attr->id = CU_LAUNCH_ATTRIBUTE_PROGRAMMATIC_STREAM_SERIALIZATION; attr->value.programmaticStreamSerializationAllowed = Py_IsTrue(kwarg); @@ -5458,16 +5467,16 @@ static PyObject* launch_impl(PyObject* const* args, Py_ssize_t nargs, CUlaunchAttribute launch_attrs[kMaxCUlaunchAttrs]; + Result driver = get_driver_api(lock); + if (!driver.is_ok()) return nullptr; + const auto num_attrs = with_block ? parse_lang_launch_kwargs(args, nargs, kwargs, launch_attrs) - : parse_tile_launch_kwargs(args, nargs, kwargs, launch_attrs); + : parse_tile_launch_kwargs(*driver, args, nargs, kwargs, launch_attrs); if (!num_attrs.is_ok()) return nullptr; - Result driver = get_driver_api(lock); - if (!driver.is_ok()) return nullptr; - if (!launch(*driver, launch_args.dispatcher, launch_args.grid, launch_args.block, launch_args.stream, launch_attrs, *num_attrs, launch_args.kernel_args, launch_args.num_kernel_args, lock)) diff --git a/changelog.d/allow-oversized-shared-memory.md b/changelog.d/allow-oversized-shared-memory.md new file mode 100644 index 00000000..74185fb6 --- /dev/null +++ b/changelog.d/allow-oversized-shared-memory.md @@ -0,0 +1,2 @@ +- Automatically allow oversized shared memory for cuTile kernel launches when run + on a driver supporting CUDA 13.4+. From e2c0ece124311963d4ea8fc34ce7d94ef2befc6b Mon Sep 17 00:00:00 2001 From: Gideon Kassa Date: Thu, 8 Oct 2026 13:34:42 -0700 Subject: [PATCH 7/7] * Fix ceiling division formula Signed-off-by: Gideon Kassa --- changelog.d/bugfix-cdiv.md | 2 ++ src/cuda/tile/_ir/ops_utils.py | 2 +- src/cuda/tile/_stub.py | 2 +- 3 files changed, 4 insertions(+), 2 deletions(-) create mode 100644 changelog.d/bugfix-cdiv.md diff --git a/changelog.d/bugfix-cdiv.md b/changelog.d/bugfix-cdiv.md new file mode 100644 index 00000000..1ce90752 --- /dev/null +++ b/changelog.d/bugfix-cdiv.md @@ -0,0 +1,2 @@ +Fixed bug in ceiling division calculations. It used to be computed using `(x + y - 1) // y` and +`(x - 1) // y + 1`, which give incorrect values for when both x and y are negative numbers. \ No newline at end of file diff --git a/src/cuda/tile/_ir/ops_utils.py b/src/cuda/tile/_ir/ops_utils.py index 8f27045f..c6296031 100644 --- a/src/cuda/tile/_ir/ops_utils.py +++ b/src/cuda/tile/_ir/ops_utils.py @@ -44,7 +44,7 @@ class MathOpDef: "sub": MathOpDef(lambda x, y: x - y, _RD_BASIC, support_flush_to_zero=True), "mul": MathOpDef(lambda x, y: x * y, _RD_BASIC, support_flush_to_zero=True), "floordiv": MathOpDef(lambda x, y: x // y), - "cdiv": MathOpDef(lambda x, y: (x + y - 1) // y), + "cdiv": MathOpDef(lambda x, y: -((-x) // y)), "truediv": MathOpDef(lambda x, y: x / y, _RD_TRUEDIV, support_flush_to_zero=True), "mod": MathOpDef(lambda x, y: x % y), "pow": MathOpDef(lambda x, y: x ** y), diff --git a/src/cuda/tile/_stub.py b/src/cuda/tile/_stub.py index b536beb0..a7a3f56d 100644 --- a/src/cuda/tile/_stub.py +++ b/src/cuda/tile/_stub.py @@ -3694,7 +3694,7 @@ def cdiv(x, y, /) -> TileOrScalar: 3 2 """ - return (x - 1) // y + 1 + return -((-x) // y) # ======== Comparison ==============