Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cuda_core/cuda/core/_launch_config.pxd
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ cdef class LaunchConfig:
public int shmem_size
public bint is_cooperative
public bint programmatic_stream_serialization
public object synchronization_policy

vector[cydriver.CUlaunchAttribute] _attrs
object __weakref__
Expand Down
9 changes: 7 additions & 2 deletions cuda_core/cuda/core/_launch_config.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -39,9 +39,12 @@ class LaunchConfig:
Whether to allow programmatic stream serialization (PDL). When True,
the kernel may overlap with a previous kernel in the same stream that
signals completion via programmatic means.
synchronization_policy : SynchronizationPolicyType | None, optional
CPU wait policy applied when synchronizing the launch stream after this
kernel. Maps to ``CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY``.
"""

def __init__(self, grid: int | tuple[int, ...] | None=None, cluster: int | tuple[int, ...] | None=None, block: int | tuple[int, ...] | None=None, shmem_size: int | None=None, is_cooperative: bool=False, programmatic_stream_serialization: bool=False) -> None:
def __init__(self, grid: int | tuple[int, ...] | None=None, cluster: int | tuple[int, ...] | None=None, block: int | tuple[int, ...] | None=None, shmem_size: int | None=None, is_cooperative: bool=False, programmatic_stream_serialization: bool=False, synchronization_policy: object=None) -> None:
"""Initialize LaunchConfig with validation.

Parameters
Expand All @@ -58,6 +61,8 @@ class LaunchConfig:
Whether to launch as cooperative kernel (default: False)
programmatic_stream_serialization : bool, optional
Whether to allow programmatic stream serialization / PDL (default: False)
synchronization_policy : SynchronizationPolicyType | None, optional
CPU wait policy for synchronizing the launch stream (default: None)
"""

def _identity(self) -> tuple[Any, ...]:
Expand All @@ -71,7 +76,7 @@ class LaunchConfig:

def __hash__(self) -> int:
...
_LAUNCH_CONFIG_ATTRS = ('grid', 'cluster', 'block', 'shmem_size', 'is_cooperative', 'programmatic_stream_serialization')
_LAUNCH_CONFIG_ATTRS = ('grid', 'cluster', 'block', 'shmem_size', 'is_cooperative', 'programmatic_stream_serialization', 'synchronization_policy')
__all__ = ['LaunchConfig']

def _to_native_launch_config(config: LaunchConfig) -> object:
Expand Down
44 changes: 44 additions & 0 deletions cuda_core/cuda/core/_launch_config.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,32 @@ _LAUNCH_CONFIG_ATTRS = (
'shmem_size',
'is_cooperative',
'programmatic_stream_serialization',
'synchronization_policy',
)


cdef object _validate_synchronization_policy(object policy):
from cuda.core.typing import SynchronizationPolicyType

if policy is None:
return None
if isinstance(policy, SynchronizationPolicyType):
return policy
try:
value = int(policy)
except (TypeError, ValueError) as exc:
raise TypeError(
"LaunchConfig.synchronization_policy must be a SynchronizationPolicyType, "
f"cuda.bindings.driver.CUsynchronizationPolicy, or int; got {type(policy).__name__}"
) from exc
try:
return SynchronizationPolicyType(value)
except ValueError as exc:
raise ValueError(
f"LaunchConfig.synchronization_policy must be one of "
f"{[member.name for member in SynchronizationPolicyType]}; got {policy!r}"
) from exc

__all__ = ['LaunchConfig']


Expand Down Expand Up @@ -59,6 +83,9 @@ cdef class LaunchConfig:
Whether to allow programmatic stream serialization (PDL). When True,
the kernel may overlap with a previous kernel in the same stream that
signals completion via programmatic means.
synchronization_policy : SynchronizationPolicyType | None, optional
CPU wait policy applied when synchronizing the launch stream after this
kernel. Maps to ``CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY``.
"""

# TODO: expand LaunchConfig to include other attributes
Expand All @@ -72,6 +99,7 @@ cdef class LaunchConfig:
shmem_size: int | None = None,
is_cooperative: bool = False,
programmatic_stream_serialization: bool = False,
synchronization_policy: object = None,
) -> None:
"""Initialize LaunchConfig with validation.

Expand All @@ -89,6 +117,8 @@ cdef class LaunchConfig:
Whether to launch as cooperative kernel (default: False)
programmatic_stream_serialization : bool, optional
Whether to allow programmatic stream serialization / PDL (default: False)
synchronization_policy : SynchronizationPolicyType | None, optional
CPU wait policy for synchronizing the launch stream (default: None)
"""
# Convert and validate grid and block dimensions
self.grid = cast_to_3_tuple("LaunchConfig.grid", grid)
Expand Down Expand Up @@ -116,6 +146,7 @@ cdef class LaunchConfig:

self.is_cooperative = is_cooperative
self.programmatic_stream_serialization = programmatic_stream_serialization
self.synchronization_policy = _validate_synchronization_policy(synchronization_policy)

if self.is_cooperative and not Device().properties.cooperative_launch:
raise CUDAError("cooperative kernels are not supported on this device")
Expand All @@ -139,6 +170,7 @@ cdef class LaunchConfig:
cdef cydriver.CUlaunchConfig _to_native_launch_config(self):
cdef cydriver.CUlaunchConfig drv_cfg
cdef cydriver.CUlaunchAttribute attr
cdef int sync_policy_value
memset(&drv_cfg, 0, sizeof(drv_cfg))
self._attrs.resize(0)

Expand Down Expand Up @@ -169,6 +201,12 @@ cdef class LaunchConfig:
attr.value.programmaticStreamSerializationAllowed = 1
self._attrs.push_back(attr)

if self.synchronization_policy is not None:
sync_policy_value = int(self.synchronization_policy)
attr.id = cydriver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY
attr.value.syncPolicy = <cydriver.CUsynchronizationPolicy>sync_policy_value
self._attrs.push_back(attr)

drv_cfg.numAttrs = self._attrs.size()
drv_cfg.attrs = self._attrs.data()

Expand Down Expand Up @@ -230,6 +268,12 @@ cpdef object _to_native_launch_config(LaunchConfig config):
attr.value.programmaticStreamSerializationAllowed = 1
attrs.append(attr)

if config.synchronization_policy is not None:
attr = driver.CUlaunchAttribute()
attr.id = driver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY
attr.value.syncPolicy = int(config.synchronization_policy)
attrs.append(attr)

drv_cfg.numAttrs = len(attrs)
drv_cfg.attrs = attrs

Expand Down
20 changes: 20 additions & 0 deletions cuda_core/cuda/core/typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
"""Public type aliases, protocols, and enumerations used in cuda.core API signatures."""

import sys
from enum import IntEnum
from typing import TYPE_CHECKING
from typing import Literal as _Literal
from typing import TypeAlias as _TypeAlias
Expand Down Expand Up @@ -47,6 +48,7 @@ class StrEnum(str, Enum):
"ProcessStateType",
"ReadModeType",
"SourceCodeType",
"SynchronizationPolicyType",
"VirtualMemoryAccessType",
"VirtualMemoryAllocationType",
"VirtualMemoryGranularityType",
Expand Down Expand Up @@ -124,6 +126,24 @@ class PCHStatusType(StrEnum):
FAILED = "failed"


class SynchronizationPolicyType(IntEnum):
"""CPU wait policy for host-side stream synchronization after a launch.

Maps to ``CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY`` and
``cuda.bindings.driver.CUsynchronizationPolicy``.

* ``AUTO`` — inherit the stream's synchronization policy.
* ``SPIN`` — busy-wait on the CPU (lowest latency).
* ``YIELD`` — yield the CPU while waiting.
* ``BLOCKING_SYNC`` — block in the OS scheduler while waiting.
"""

AUTO = driver.CUsynchronizationPolicy.CU_SYNC_POLICY_AUTO
SPIN = driver.CUsynchronizationPolicy.CU_SYNC_POLICY_SPIN
YIELD = driver.CUsynchronizationPolicy.CU_SYNC_POLICY_YIELD
BLOCKING_SYNC = driver.CUsynchronizationPolicy.CU_SYNC_POLICY_BLOCKING_SYNC


class GraphConditionalType(StrEnum):
"""Conditional node flavor for :class:`~cuda.core.graph.GraphBuilder`.

Expand Down
111 changes: 110 additions & 1 deletion cuda_core/tests/test_launcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
import pytest
from conftest import skipif_need_cuda_headers

from cuda.bindings import driver
from cuda.core import (
Device,
DeviceMemoryResource,
Expand All @@ -26,7 +27,7 @@
)
from cuda.core._memory._legacy import _SynchronousMemoryResource
from cuda.core._utils.cuda_utils import CUDAError
from cuda.core.typing import ObjectCodeFormatType, SourceCodeType
from cuda.core.typing import ObjectCodeFormatType, SourceCodeType, SynchronizationPolicyType


def test_launch_config_init(init_cuda):
Expand Down Expand Up @@ -202,6 +203,114 @@ def test_to_native_launch_config_pdl():
)


@pytest.mark.parametrize(
("policy", "expected_value"),
[
(SynchronizationPolicyType.AUTO, driver.CUsynchronizationPolicy.CU_SYNC_POLICY_AUTO),
(SynchronizationPolicyType.SPIN, driver.CUsynchronizationPolicy.CU_SYNC_POLICY_SPIN),
(SynchronizationPolicyType.YIELD, driver.CUsynchronizationPolicy.CU_SYNC_POLICY_YIELD),
(
SynchronizationPolicyType.BLOCKING_SYNC,
driver.CUsynchronizationPolicy.CU_SYNC_POLICY_BLOCKING_SYNC,
),
(driver.CUsynchronizationPolicy.CU_SYNC_POLICY_SPIN, driver.CUsynchronizationPolicy.CU_SYNC_POLICY_SPIN),
],
)
def test_to_native_launch_config_synchronization_policy(policy, expected_value):
"""LaunchConfig.synchronization_policy maps to CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY."""
from cuda.core._launch_config import _to_native_launch_config

config = LaunchConfig(grid=1, block=1, synchronization_policy=policy)
assert config.synchronization_policy == SynchronizationPolicyType(int(expected_value))

native = _to_native_launch_config(config)
assert native.numAttrs == 1
attr = native.attrs[0]
assert attr.id == driver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY
assert attr.value.syncPolicy == expected_value


def test_launch_config_synchronization_policy_default():
config = LaunchConfig(grid=1, block=1)
assert config.synchronization_policy is None

from cuda.core._launch_config import _to_native_launch_config

native = _to_native_launch_config(config)
assert native.numAttrs == 0


@pytest.mark.parametrize("invalid_policy", ["spin", -1, 99])
def test_launch_config_synchronization_policy_invalid(invalid_policy):
with pytest.raises((TypeError, ValueError)):
LaunchConfig(grid=1, block=1, synchronization_policy=invalid_policy)


def test_to_native_launch_config_synchronization_policy_with_cooperative(monkeypatch):
"""synchronization_policy can be combined with other launch attributes."""
from cuda.core import _launch_config as _lc_mod
from cuda.core._launch_config import _to_native_launch_config

class _FakeProps:
cooperative_launch = True

class _FakeDev:
properties = _FakeProps()

monkeypatch.setattr(_lc_mod, "Device", lambda: _FakeDev())

config = LaunchConfig(
grid=1,
block=1,
is_cooperative=True,
synchronization_policy=SynchronizationPolicyType.SPIN,
)
native = _to_native_launch_config(config)
assert native.numAttrs == 2
attr_ids = {attr.id for attr in native.attrs}
assert driver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_COOPERATIVE in attr_ids
assert driver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY in attr_ids
sync_attrs = [
attr
for attr in native.attrs
if attr.id == driver.CUlaunchAttributeID.CU_LAUNCH_ATTRIBUTE_SYNCHRONIZATION_POLICY
]
assert len(sync_attrs) == 1
assert sync_attrs[0].value.syncPolicy == driver.CUsynchronizationPolicy.CU_SYNC_POLICY_SPIN


@pytest.mark.parametrize(
"policy",
[
SynchronizationPolicyType.AUTO,
SynchronizationPolicyType.SPIN,
SynchronizationPolicyType.YIELD,
SynchronizationPolicyType.BLOCKING_SYNC,
],
)
def test_launch_with_synchronization_policy(init_cuda, policy):
"""Driver accepts per-launch synchronization policies on a real kernel launch."""
import cuda.bindings
from cuda.core._utils.version import driver_version

if int(cuda.bindings.__version__.split(".")[0]) >= 13 and driver_version()[0] < 13:
pytest.skip("CUDA 13 bindings produce modules incompatible with CUDA 12 drivers")

dev = Device()
dev.set_current()
stream = dev.create_stream()

code = 'extern "C" __global__ void noop() {}'
arch = "".join(f"{i}" for i in dev.compute_capability)
program = Program(code, SourceCodeType.CXX, options=ProgramOptions(arch=f"sm_{arch}"))
mod = program.compile(ObjectCodeFormatType.CUBIN)
ker = mod.get_kernel("noop")

config = LaunchConfig(grid=1, block=1, synchronization_policy=policy)
launch(stream, config, ker)
stream.sync()


@skipif_need_cuda_headers
def test_pdl_primary_secondary_overlap_same_stream():
"""Primary + secondary PDL launch on one stream can overlap on Hopper+.
Expand Down
2 changes: 1 addition & 1 deletion cuda_core/tests/test_object_protocols.py
Original file line number Diff line number Diff line change
Expand Up @@ -685,7 +685,7 @@ def sample_switch_node_alt(sample_graphdef):
"sample_launch_config",
r"LaunchConfig\(grid=\(\d+, \d+, \d+\), cluster=.+, block=\(\d+, \d+, \d+\), "
r"shmem_size=\d+, is_cooperative=(?:True|False), "
r"programmatic_stream_serialization=(?:True|False)\)",
r"programmatic_stream_serialization=(?:True|False), synchronization_policy=(?:None|SynchronizationPolicyType\.\w+)\)",
),
("sample_kernel", r"<Kernel handle=0x[0-9a-f]+>"),
# ObjectCode variations (by code_type)
Expand Down
Loading