Rate this Page

Source code for helion.language.atomic_ops

from __future__ import annotations

import ast
import itertools
from typing import Callable

import torch
from torch._inductor.codegen.simd import constant_repr
from torch.fx import has_side_effect

from .. import exc
from .._compiler.ast_extension import expr_from_string
from .._compiler.indexing_strategy import SubscriptIndexing
from . import _decorators

__all__ = [
    "atomic_add",
    "atomic_and",
    "atomic_cas",
    "atomic_max",
    "atomic_min",
    "atomic_or",
    "atomic_xchg",
    "atomic_xor",
]


_VALID_SEMS: set[str] = {"relaxed", "acquire", "release", "acq_rel"}


def _validate_sem(sem: str) -> None:
    if sem not in _VALID_SEMS:
        raise exc.InternalError(
            ValueError(
                f"Invalid memory semantic '{sem}'. Valid options are: relaxed, acquire, release, acq_rel"
            )
        )


def _prepare_mem_args(
    target: torch.Tensor,
    index: list[object],
    *values: object,
    sem: str = "relaxed",
) -> tuple:
    from .tile_proxy import Tile

    _validate_sem(sem)
    index = Tile._prepare_index(index)
    index = Tile._tiles_to_sizes_for_index(index)
    return (target, index, *values, sem)


def _to_ast_values(values: list[object]) -> list[ast.AST]:
    out: list[ast.AST] = []
    for v in values:
        if isinstance(v, (int, float, bool)):
            out.append(expr_from_string(constant_repr(v)))
        else:
            assert isinstance(v, ast.AST)
            out.append(v)
    return out


def _ref_atomic_binop(
    target: torch.Tensor,
    index: list[object],
    value: torch.Tensor | float | bool,
    op: Callable[[torch.Tensor, torch.Tensor], torch.Tensor],
) -> torch.Tensor:
    """Shared ref implementation for simple atomic binary ops (xchg/and/or/xor/max/min).

    Processes indices, clones the previous value, applies the op, and returns prev.
    For xchg, pass op=lambda old, val: val.
    """
    from .ref_tile import RefTile

    processed_index: list[object] = []
    for idx in index:
        if isinstance(idx, RefTile):
            processed_index.append(idx._slice)
        elif isinstance(idx, torch.Tensor) and idx.numel() == 1:
            processed_index.append(int(idx.item()))
        else:
            processed_index.append(idx)
    idx_tuple = tuple(processed_index)
    # pyrefly: ignore [bad-index]
    prev = target[idx_tuple].clone()
    val = (
        value
        if isinstance(value, torch.Tensor)
        else torch.as_tensor(value, dtype=target.dtype, device=target.device)
    )
    # pyrefly: ignore [bad-index, unsupported-operation]
    target[idx_tuple] = op(target[idx_tuple], val)
    return prev


def _ref_apply(
    target: torch.Tensor,
    index: list[object],
    apply_fn: Callable[[torch.Tensor, tuple, object], None],
    value: object,
) -> None:
    from .ref_tile import RefTile

    # Convert indices to proper format
    processed_index: list[object] = []
    for idx in index:
        if isinstance(idx, RefTile):
            processed_index.append(idx._slice)
        elif isinstance(idx, torch.Tensor) and idx.numel() == 1:
            processed_index.append(int(idx.item()))
        else:
            processed_index.append(idx)

    # Find tensor indices that need element-wise processing
    tensor_indices = [
        (i, idx)
        for i, idx in enumerate(processed_index)
        if isinstance(idx, torch.Tensor) and idx.numel() > 1
    ]

    if tensor_indices:
        # Element-wise processing for tensor indices (handle first tensor index)
        i, tensor_idx = tensor_indices[0]

        if tensor_idx.ndim == 0:
            coords_iter = [()]
        else:
            ranges = [range(dim) for dim in tensor_idx.shape]
            coords_iter = itertools.product(*ranges)

        for coords in coords_iter:
            elem = tensor_idx[coords].item()
            new_index = processed_index.copy()
            new_index[i] = int(elem)
            if isinstance(value, torch.Tensor) and value.numel() > 1:
                next_value = value[coords]
            else:
                next_value = value
            _ref_apply(target, new_index, apply_fn, next_value)
    else:
        apply_fn(target, tuple(processed_index), value)


# -- atomic_add --


[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_add( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically add a value to a target tensor. Performs an atomic read-modify-write that adds ``value`` to ``target[index]``. This is safe for concurrent access from multiple threads/blocks. Args: target: Tensor to update. index: Indices selecting elements to update. Can include tiles. value: Value(s) to add (tensor or scalar). sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. Example: @helion.kernel def global_sum(x: torch.Tensor, result: torch.Tensor) -> torch.Tensor: for tile in hl.tile(x.size(0)): hl.atomic_add(result, [0], x[tile].sum()) return result Notes: - Use for race-free accumulation across parallel execution. - Higher memory semantics may reduce performance. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_add) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> tuple[torch.Tensor, object, torch.Tensor | float | int, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_add) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_add) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) from .ref_tile import RefTile # Convert indices for shape computation and fast path detection processed_index: list[object] = [] has_tensor_index = False for idx in index: if isinstance(idx, RefTile): processed_index.append(idx._slice) elif isinstance(idx, torch.Tensor): if idx.numel() == 1: processed_index.append(int(idx.item())) else: processed_index.append(idx) has_tensor_index = True else: processed_index.append(idx) def _convert_value_to_target_dtype(val: object) -> torch.Tensor: if isinstance(val, torch.Tensor): vt = val.to(device=target.device) if vt.dtype != target.dtype: vt = vt.to(dtype=target.dtype) return vt return torch.as_tensor(val, dtype=target.dtype, device=target.device) if has_tensor_index: ret_shape = SubscriptIndexing.compute_shape(target, processed_index) prev_chunks: list[torch.Tensor] = [] def apply(t: torch.Tensor, idx_tuple: tuple, v: object) -> None: prev_val = t[idx_tuple].clone() val_tensor = _convert_value_to_target_dtype(v) t[idx_tuple] = t[idx_tuple] + val_tensor prev_chunks.append(prev_val.reshape(-1)) _ref_apply(target, index, apply, value) if prev_chunks: flat_prev = torch.cat(prev_chunks) else: flat_prev = target.new_empty(0, dtype=target.dtype, device=target.device) return flat_prev.reshape(ret_shape) idx_tuple = tuple(processed_index) # pyrefly: ignore [bad-index] prev = target[idx_tuple].clone() val_tensor = _convert_value_to_target_dtype(value) # pyrefly: ignore [bad-index, unsupported-operation] target[idx_tuple] = target[idx_tuple] + val_tensor return prev # -- atomic_xchg --
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_xchg( target: torch.Tensor, index: list[object], value: torch.Tensor | float | bool, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically exchange (set) a value at ``target[index]``. Args: target: Tensor to update. index: Indices selecting elements to update. Can include tiles. value: New value(s) to set. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_xchg) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float | bool, sem: str = "relaxed", ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_xchg) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_xchg) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float | bool, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, lambda old, val: val) # -- atomic_and/or/xor --
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_and( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically apply bitwise AND with ``value`` to ``target[index]``. Args: target: Tensor to update (integer/bool dtype). index: Indices selecting elements to update. Can include tiles. value: Value(s) to AND with. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_and) def _( target: torch.Tensor, index: list[object], value: object, sem: str = "relaxed" ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_and) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_and) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, torch.bitwise_and)
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_or( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically apply bitwise OR with ``value`` to ``target[index]``. Args: target: Tensor to update (integer/bool dtype). index: Indices selecting elements to update. Can include tiles. value: Value(s) to OR with. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_or) def _( target: torch.Tensor, index: list[object], value: object, sem: str = "relaxed" ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_or) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_or) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, torch.bitwise_or)
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_xor( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically apply bitwise XOR with ``value`` to ``target[index]``. Args: target: Tensor to update (integer/bool dtype). index: Indices selecting elements to update. Can include tiles. value: Value(s) to XOR with. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_xor) def _( target: torch.Tensor, index: list[object], value: object, sem: str = "relaxed" ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_xor) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_xor) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | int | bool, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, torch.bitwise_xor) # -- atomic_max/min --
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_max( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically update ``target[index]`` with the maximum of current value and ``value``. Args: target: Tensor to update. index: Indices selecting elements to update. Can include tiles. value: Value(s) to compare with. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_max) def _( target: torch.Tensor, index: list[object], value: object, sem: str = "relaxed" ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_max) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_max) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, torch.maximum)
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_min( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically update ``target[index]`` with the minimum of current value and ``value``. Args: target: Tensor to update. index: Indices selecting elements to update. Can include tiles. value: Value(s) to compare with. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the update. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_min) def _( target: torch.Tensor, index: list[object], value: object, sem: str = "relaxed" ) -> tuple[torch.Tensor, object, object, str]: return _prepare_mem_args(target, index, value, sem=sem) @_decorators.register_fake(atomic_min) def _( target: torch.Tensor, index: list[object], value: torch.Tensor, sem: str = "relaxed" ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_min) def _( target: torch.Tensor, index: list[object], value: torch.Tensor | float, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) return _ref_atomic_binop(target, index, value, torch.minimum) # -- atomic_cas --
[docs] @has_side_effect @_decorators.api(allow_host_tensor=True, tiles_as_sizes=True) def atomic_cas( target: torch.Tensor, index: list[object], expected: torch.Tensor | float | bool, value: torch.Tensor | float | bool, sem: str = "relaxed", ) -> torch.Tensor: """ Atomically compare-and-swap a value at ``target[index]``. If the current value equals ``expected``, writes ``value``. Otherwise leaves memory unchanged. Args: target: Tensor to update. index: Indices selecting elements to update. Can include tiles. expected: Expected current value(s) used for comparison. value: New value(s) to write if comparison succeeds. sem: Memory ordering semantics. One of ``"relaxed"``, ``"acquire"``, ``"release"``, ``"acq_rel"``. Defaults to ``"relaxed"``. Returns: torch.Tensor: The previous value(s) stored at ``target[index]`` before the compare-and-swap. Note: Triton CAS doesn’t support a masked form; our generated code uses an unmasked CAS and relies on index masking to avoid OOB. """ raise exc.NotInsideKernel
@_decorators.prepare_args(atomic_cas) def _( target: torch.Tensor, index: list[object], expected: object, value: object, sem: str = "relaxed", ) -> tuple[torch.Tensor, object, object, object, str]: return _prepare_mem_args(target, index, expected, value, sem=sem) @_decorators.register_fake(atomic_cas) def _( target: torch.Tensor, index: list[object], expected: torch.Tensor, value: torch.Tensor, sem: str = "relaxed", ) -> torch.Tensor: target_shape = SubscriptIndexing.compute_shape(target, index) return target.new_empty(target_shape) @_decorators.ref(atomic_cas) def _( target: torch.Tensor, index: list[object], expected: torch.Tensor | float | bool, value: torch.Tensor | float | bool, sem: str = "relaxed", ) -> torch.Tensor: _validate_sem(sem) from .ref_tile import RefTile processed_index: list[object] = [] for idx in index: if isinstance(idx, RefTile): processed_index.append(idx._slice) elif isinstance(idx, torch.Tensor) and idx.numel() == 1: processed_index.append(int(idx.item())) else: processed_index.append(idx) idx_tuple = tuple(processed_index) # pyrefly: ignore [bad-index] prev = target[idx_tuple].clone() exp_t = ( expected if isinstance(expected, torch.Tensor) else torch.as_tensor(expected, dtype=target.dtype, device=target.device) ) val_t = ( value if isinstance(value, torch.Tensor) else torch.as_tensor(value, dtype=target.dtype, device=target.device) ) # pyrefly: ignore [bad-index] mask = target[idx_tuple] == exp_t # pyrefly: ignore [bad-index, unsupported-operation] target[idx_tuple] = torch.where(mask, val_t, target[idx_tuple]) return prev ATOMIC_OPS: frozenset[Callable[..., object]] = frozenset( { atomic_add, atomic_and, atomic_cas, atomic_max, atomic_min, atomic_or, atomic_xchg, atomic_xor, } )