tilelang.ascend.language.simd ============================= .. py:module:: tilelang.ascend.language.simd .. autoapi-nested-parse:: T.simd.* - Raw CCE vector intrinsics for Ascend SIMD programming. Function names match the underlying CCE intrinsics directly. MODE_MERGING preserves inactive lanes of the mutable destination register. On Ascend 950, validated 8/16/32-bit operations use the CCE merging overloads; ``vdupv`` maps to CCE's vector ``vdup`` overload. Scalar BF16 ``vdup`` retains software merging to work around CANN 9.2's inactive-lane bug. Precision-specific SFU algorithms also retain their wrappers, including the default exact FP32 division and the ``ftz_false`` variants of ``vexp``, ``vln``, and ``vsqrt``. Classes ------- .. autoapisummary:: tilelang.ascend.language.simd.SimdPair Functions --------- .. autoapisummary:: tilelang.ascend.language.simd.pset tilelang.ascend.language.simd.pge tilelang.ascend.language.simd.update_mask tilelang.ascend.language.simd.pand tilelang.ascend.language.simd.por tilelang.ascend.language.simd.pxor tilelang.ascend.language.simd.pnot tilelang.ascend.language.simd.psel tilelang.ascend.language.simd.ppack tilelang.ascend.language.simd.punpack tilelang.ascend.language.simd.pintlv tilelang.ascend.language.simd.pdintlv tilelang.ascend.language.simd.alloc_var tilelang.ascend.language.simd.alloc_local tilelang.ascend.language.simd.pld tilelang.ascend.language.simd.pst tilelang.ascend.language.simd.vld tilelang.ascend.language.simd.vld2 tilelang.ascend.language.simd.vsts tilelang.ascend.language.simd.vsstb tilelang.ascend.language.simd.make_ubuf_ptr tilelang.ascend.language.simd.vadd tilelang.ascend.language.simd.vaddc tilelang.ascend.language.simd.vsubc tilelang.ascend.language.simd.vaddcs tilelang.ascend.language.simd.vsubcs tilelang.ascend.language.simd.vmull tilelang.ascend.language.simd.vsub tilelang.ascend.language.simd.vmul tilelang.ascend.language.simd.vmula tilelang.ascend.language.simd.vmadd tilelang.ascend.language.simd.vaxpy tilelang.ascend.language.simd.dhistv2 tilelang.ascend.language.simd.chistv2 tilelang.ascend.language.simd.vdiv tilelang.ascend.language.simd.vmax tilelang.ascend.language.simd.vmin tilelang.ascend.language.simd.vand tilelang.ascend.language.simd.vor tilelang.ascend.language.simd.vxor tilelang.ascend.language.simd.vshl tilelang.ascend.language.simd.vshr tilelang.ascend.language.simd.vexp tilelang.ascend.language.simd.vln tilelang.ascend.language.simd.vsqrt tilelang.ascend.language.simd.vabs tilelang.ascend.language.simd.vneg tilelang.ascend.language.simd.vrelu tilelang.ascend.language.simd.vlrelu tilelang.ascend.language.simd.vprelu tilelang.ascend.language.simd.vnot tilelang.ascend.language.simd.vdup tilelang.ascend.language.simd.vdupv tilelang.ascend.language.simd.vcpadd tilelang.ascend.language.simd.vcadd tilelang.ascend.language.simd.vcmax tilelang.ascend.language.simd.vcmin tilelang.ascend.language.simd.vcgadd tilelang.ascend.language.simd.vcgmax tilelang.ascend.language.simd.vcgmin tilelang.ascend.language.simd.vsqz tilelang.ascend.language.simd.vusqz tilelang.ascend.language.simd.vci tilelang.ascend.language.simd.vcmp tilelang.ascend.language.simd.vcmps tilelang.ascend.language.simd.vintlv tilelang.ascend.language.simd.vdintlv tilelang.ascend.language.simd.pair_get tilelang.ascend.language.simd.vpack tilelang.ascend.language.simd.vunpack tilelang.ascend.language.simd.vgatherb tilelang.ascend.language.simd.vgather2 tilelang.ascend.language.simd.vscatter tilelang.ascend.language.simd.vexpdif tilelang.ascend.language.simd.vabsdif tilelang.ascend.language.simd.vcvt tilelang.ascend.language.simd.vsel tilelang.ascend.language.simd.vselr tilelang.ascend.language.simd.vmaxs tilelang.ascend.language.simd.vmins tilelang.ascend.language.simd.vmuls tilelang.ascend.language.simd.vadds tilelang.ascend.language.simd.vshls tilelang.ascend.language.simd.vshrs tilelang.ascend.language.simd.mem_bar Module Contents --------------- .. py:class:: SimdPair(pair, dtype=None) Wraps a two-result SIMD intrinsic and preserves each result dtype. ``a, b = ...`` emits two ``pair_get`` calls against the pair. The pair- producing op (e.g. ``vintlv``, ``vld2``, or post-update ``vld``) is bound once at its call site by the frontend, so both ``pair_get`` calls reference a single bound variable rather than inlining the pair expression twice. ``dtype`` is a pair describing both result types, such as the ``(boolx256, int32x64)`` carry/result pair returned by ``vaddc``. Omitting it defaults both result types to the dtype of the backing TIR expression. .. py:property:: dtype .. py:method:: __getitem__(index) .. py:method:: __iter__() .. py:function:: pset(elem_width, dist = 'PAT_ALL') Create a predicate mask: pset_bXX(dist). Returns a vector_bool (boolx256). dist: "PAT_ALL", "PAT_VL1".."PAT_VL128", "PAT_M3", "PAT_M4", "PAT_H", "PAT_Q", etc. .. py:function:: pge(elem_width, dist = 'PAT_ALL') Create a predicate mask from pge_bXX(dist). .. py:function:: update_mask(value, width=32) Runtime tail predicate: lanes [0, value) active (b8/b16/b32). .. py:function:: pand(src0, src1, mask) .. py:function:: por(src0, src1, mask) .. py:function:: pxor(src0, src1, mask) .. py:function:: pnot(src, mask) .. py:function:: psel(src0, src1, mask) .. py:function:: ppack(src, part=0) Predicate pack 2:1 (zeroing): dst = ppack(src, LOWER/HIGHER). .. py:function:: punpack(src, part=0) Predicate unpack 1:2 (zeroing): dst = punpack(src, LOWER/HIGHER). .. py:function:: pintlv(src0, src1, width=32) Predicate interleave -> pair of predicates (b8/b16/b32). .. py:function:: pdintlv(src0, src1, width=32) Predicate deinterleave -> pair of predicates (b8/b16/b32). .. py:function:: alloc_var(dtype) Allocate a single mutable SIMD register variable (return-value style). Uses ``local.var`` scope and behaves as one vector register value. For an addressable array of registers (``v[i]``), use :func:`alloc_local`. .. py:function:: alloc_local(shape, dtype) Allocate an addressable array of mutable SIMD register variables. Uses ``local`` scope so each element is an individually addressable register, allowing indexed access like ``v[i]``:: v = T.simd.alloc_local(4, "float32") for i in T.Unroll(4, explicit=True): v[i] = vld(s_ub[i * VL]) .. py:function:: pld(addr, dist='NORM') Load a predicate from UB. Express address offsets in ``addr``. .. py:function:: pst(addr, src, dist='NORM') Store a predicate to UB. Express address offsets in ``addr``. .. py:function:: vld(addr, dist='NORM', *, post_inc=None) Vector load. Returns a typed vector register. addr can be a BufferLoad auto-wrapped as tl.access_ptr. Express address offsets in ``addr``. With ``post_inc=step``, load through a mutable :func:`make_ubuf_ptr` handle and return ``(vector, advanced_pointer)``. Assign the second result back to the handle. ``step`` is a signed int32 increment in elements of the dtype declared by :func:`make_ubuf_ptr`; the load uses the old address. The dtype must be 8/16/32-bit and match the distribution. ``None`` selects an ordinary load; zero still returns the pair without advancing the pointer:: src_ptr = T.simd.make_ubuf_ptr(T.access_ptr(src_ub[0], "r", extent=256), "uint16") first, src_ptr = T.simd.vld(src_ptr, post_inc=128) second, src_ptr = T.simd.vld(src_ptr, post_inc=128) Keep the pointer within one SIMD VF and declare its complete accessed span in the initializer's ``access_ptr``, as in the example. A ``BRC_B8/B16/B32`` broadcast replicates one element of the width named by the suffix, widening the result past the source buffer's element type when the two differ (``BRC_B16`` over a ``uint8`` buffer broadcasts a 16-bit element, not a byte). Every other distribution takes its element width from the source buffer; their ``_B*`` suffixes describe the data being loaded and must agree with it. .. py:function:: vld2(addr, dist='DINTLV_B16', off=None) Dual-dest vector load: ``a, b = vld2(x_ub[i, col], dist="DINTLV_B8")``. Supported dists: - ``DINTLV_B8``: load 512xu8/fp8 -> two 256-lane regs (even/odd bytes) - ``DINTLV_B16``: load 256xbf16/u16 -> two 128-lane regs - ``DINTLV_B32``: load 128xf32/u32 -> two 64-lane regs The access_ptr footprint is always 2x the single-vector width. This is an opaque memory load: the pair is bound once at this program point so the two ``pair_get`` calls from the unpack share a single load rather than issuing two independent loads. .. py:function:: vsts(addr, src, mask=None, dist='NORM_B32', extent=None) Vector store to ``addr``. addr can be a BufferLoad auto-wrapped as tl.access_ptr. Express address offsets in ``addr``. Optional ``extent`` overrides the default access_ptr footprint (e.g. 8 for a dense PAT_VL8 NORM_B16 recip pack). .. py:function:: vsstb(src, base, stride, mask=None, update=False) Scatter-store 32B blocks with an optional POST_UPDATE pointer. Passing a regular buffer access performs a store and returns ``void``. Set ``update=True`` with the mutable handle returned by :func:`make_ubuf_ptr` to enable POST_UPDATE and return the advanced pointer, which should be assigned back to the same handle:: dst_ptr = T.simd.make_ubuf_ptr(dst_ub[0], "bfloat16") dst_ptr = T.simd.vsstb(src, dst_ptr, stride, mask, update=True) .. py:function:: make_ubuf_ptr(buf_access, dtype) Allocate a mutable UB pointer for post-update :func:`vld` / :func:`vsstb` calls. ``dtype`` declares the pointee element type used by :func:`vld` and checked against the source vector by :func:`vsstb`. It is stored in the enclosing block's IR annotations; the mutable carrier remains a ``handle`` buffer. The pointer is carried by a ``local.var`` handle buffer. Assigning the advanced handle returned by :func:`vld` or :func:`vsstb` writes it back into the same mutable carrier:: dst_ptr = T.simd.make_ubuf_ptr(dst_ub[0], "bfloat16") dst_ptr = T.simd.vsstb(src, dst_ptr, stride, mask, update=True) .. py:function:: vadd(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vaddc(src0, src1, mask=None) Add int32/uint32 vectors without carry-in and return ``(carry, result)``. ``carry`` is a ``boolx256`` predicate register and ``result`` has the same dtype as the inputs. The underlying Ascend ``vaddc`` instruction only supports full-register ``int32x64`` and ``uint32x64`` operands. .. py:function:: vsubc(src0, src1, mask=None) Subtract int32/uint32 without carry-in and return ``(carry, result)``. ``carry`` is 1 where the subtraction completes without borrow. .. py:function:: vaddcs(src0, src1, carrysrcp, mask=None) Add int32/uint32 with carry-in predicate and return ``(carry, result)``. .. py:function:: vsubcs(src0, src1, carrysrcp, mask=None) Subtract int32/uint32 with carry-in predicate and return ``(carry, result)``. .. py:function:: vmull(src0, src1, mask=None) Widening 32x32->64 multiply returning ``(lo, hi)`` (int32/uint32). .. py:function:: vsub(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vmul(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vmula(dst, src0, src1, mask=None, mode='MODE_ZEROING') Fused multiply-add: dst = dst + src0 * src1. .. py:function:: vmadd(dst, src0, src1, mask=None, mode='MODE_ZEROING') Fused multiply-add: dst = dst * src0 + src1. .. py:function:: vaxpy(dst, src, scalar, mask=None, mode='MODE_ZEROING') Fused scalar multiply-add: dst = src * scalar + dst. .. py:function:: dhistv2(dst, src, mask=None, bin=0) Accumulate a frequency histogram of a uint8 vector into a uint16 vector. ``bin=0`` counts values in ``[0, 127]`` and ``bin=1`` counts values in ``[128, 255]``. The destination register is updated in place. .. py:function:: chistv2(dst, src, mask=None, bin=0) Accumulate a cumulative histogram of a uint8 vector into a uint16 vector. ``bin=0`` returns cumulative counts through values ``[0, 127]`` and ``bin=1`` returns cumulative counts through values ``[128, 255]``. The destination register is updated in place. .. py:function:: vdiv(src0, src1, mask=None, mode='MODE_ZEROING', precision=None) Divide vectors, optionally selecting the SFU implementation. ``precision=None`` follows ``tl.enable_fast_math`` (precise fp32 division when fast math is off, hardware instruction otherwise). Non-fp32 division always uses the hardware instruction. ``precision`` selects the implementation per op: - ``'ftz_true'``: bare hardware SFU (flush-to-zero semantics) - ``'exact'`` (alias ``'vdiv_0ulp_ftz_true'``): correctly-rounded fp32 division (CANN DivAlgo::PRECISION_0ULP_FTZ_TRUE / DivPrecisionImpl) Requires float32 for precise paths. .. py:function:: vmax(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vmin(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vand(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vor(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vxor(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vshl(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vshr(src0, src1, mask=None, mode='MODE_ZEROING') .. py:function:: vexp(src, mask=None, mode='MODE_ZEROING', precision=None) .. py:function:: vln(src, mask=None, mode='MODE_ZEROING', precision=None) .. py:function:: vsqrt(src, mask=None, mode='MODE_ZEROING', precision=None) .. py:function:: vabs(src, mask=None, mode='MODE_ZEROING') .. py:function:: vneg(src, mask=None, mode='MODE_ZEROING') .. py:function:: vrelu(src, mask=None, mode='MODE_ZEROING') .. py:function:: vlrelu(src, alpha, mask=None) Leaky ReLU with scalar slope (f16/f32): dst = src >= 0 ? src : alpha * src. .. py:function:: vprelu(src0, src1, mask=None) Parametric ReLU with per-lane slope vector (f16/f32). .. py:function:: vnot(src, mask=None, mode='MODE_ZEROING') .. py:function:: vdup(src, dtype_str, mask=None, mode='MODE_ZEROING') Broadcast scalar to all lanes: dst = vdup(scalar, dtype_str, mask). dtype_str specifies the target vector element type (e.g. "float32"). .. py:function:: vdupv(src, mask=None, pos='POS_LOWEST', mode='MODE_ZEROING') Broadcast lane N of src vector to all lanes: dst = vdupv(src, mask, pos). pos: "POS_LOWEST" (lane 0) or "POS_HIGHEST" (lane N). .. py:function:: vcpadd(src, mask=None, mode='MODE_ZEROING') Pairwise adjacent-lane add. The sums of adjacent source lanes are packed into the low half of the result vector. Supported types: float32, float16. .. py:function:: vcadd(src, mask=None, mode='MODE_ZEROING') Pairwise add reduction: dst = vcadd(src, mask). .. py:function:: vcmax(src, mask=None, mode='MODE_ZEROING') Pairwise max reduction: dst = vcmax(src, mask). .. py:function:: vcmin(src, mask=None, mode='MODE_ZEROING') Pairwise min reduction: dst = vcmin(src, mask). .. py:function:: vcgadd(src, mask=None, mode='MODE_ZEROING') Grouped add reduction: dst = vcgadd(src, mask). .. py:function:: vcgmax(src, mask=None, mode='MODE_ZEROING') Grouped max reduction: dst = vcgmax(src, mask). .. py:function:: vcgmin(src, mask=None, mode='MODE_ZEROING') Grouped min reduction: dst = vcgmin(src, mask). .. py:function:: vsqz(src, mask=None, mode='MODE_STORED') Squeeze selected lanes toward the low lanes: dst = vsqz(src, mask). .. py:function:: vusqz(mask, dtype='int32') Per-lane exclusive prefix count of mask (s8/s16/s32). .. py:function:: vci(index, dtype_str, order='INC_ORDER') Index ramp: dst = vci(index, dtype_str, order). dst[lane] = index + lane (INC_ORDER) / index - lane (DEC_ORDER). dtype_str is the destination vector element type (e.g. "int32", "float32"). .. py:function:: vcmp(src0, src1, mask=None, op='eq') Elementwise compare -> vector_bool: dst = vcmp_(src0, src1, mask). op: eq/ne/gt/ge/lt/le. .. py:function:: vcmps(src, scalar, mask=None, op='lt') Compare vector vs scalar -> vector_bool: dst = vcmps_(src, scalar, mask). op: eq/ne/gt/ge/lt/le. .. py:function:: vintlv(src0, src1) Interleave two vector registers. Unpack: a, b = vintlv(x, y) The pair is bound once at the call site so the unpack's two ``pair_get`` calls share a single permutation instruction instead of inlining the pair expression twice. .. py:function:: vdintlv(src0, src1) De-interleave two vector registers. Unpack: a, b = vdintlv(x, y) The pair is bound once at the call site so the unpack's two ``pair_get`` calls share a single permutation instruction instead of inlining the pair expression twice. .. py:function:: pair_get(pair, index) Extract element from a pair: reg = pair_get(pair, 0) or pair_get(pair, 1). .. py:function:: vpack(src, part=0) Pack wider lanes to narrower lanes: dst = vpack(src, LOWER/HIGHER). Supports u32->u16 and u16->u8 (needed for dense UE8M0 scale packing). .. py:function:: vunpack(src, part=0) Widen half of src: u8->u16, s8->s16, u16->u32, s16->s32. .. py:function:: vgatherb(base, index, mask=None) Gather 32B blocks from base using vector_u32 block offsets. Returns a vector. .. py:function:: vgather2(base, index, mask=None) Gather elements from base using per-lane offsets. .. py:function:: vscatter(src, base, index, mask=None) Scatter-store: base[index[lane]] = src[lane]. Write-side counterpart to vgatherb/vgather2. ``index`` is a vector register of per-lane element offsets (uint32 for 32-bit elements, uint16 for 8/16-bit). .. py:function:: vexpdif(src0, src1, mask=None) Fused exp-sub: dst = exp(src0 - src1). The same-width form supports matching float32 vectors. Widening float16 inputs to float32 requires separate even/odd results and is not represented by this API. .. py:function:: vabsdif(src0, src1, mask=None, mode='MODE_ZEROING') Fused abs-sub: dst = vabsdif(src0, src1, mask, mode). Computes dst = abs(src0 - src1) in a single instruction. Supported types: float32, float16. .. py:function:: vcvt(src, target_dtype, mask=None, round='ROUND_R', sat=True, part=0, mode='MODE_ZEROING') Vector type conversion between float32, float16, bfloat16, float8, float4, and integers. :param src: Source vector register holding the input element values. :type src: VReg :param target_dtype: Destination element type, e.g. "float32", "float16", "bfloat16", "float8_e4m3", "float8_e5m2", "float4_e2m1fn", "int32", "int16", "uint16", "int8", "uint8", "int64". :type target_dtype: T.dtype :param mask: Predicate mask; only lanes where mask[i] is true participate and are written. Defaults to an all-lanes mask matching the narrower (low-lane) element: the source width when widening, the target width when narrowing. :type mask: VReg (bool) :param round: IEEE-754 rounding mode. Valid values: - ``"ROUND_R"`` (default): Round to nearest, ties to even (banker's rounding). - ``"ROUND_A"``: Round away from zero. - ``"ROUND_F"``: Round toward -inf (floor). - ``"ROUND_C"``: Round toward +inf (ceiling). - ``"ROUND_Z"``: Round toward zero (truncation). - ``"ROUND_O"``: Round to odd (only for f32->f16). - ``"ROUND_H"``: Round half away from zero (only for hif8 conversions). Not all modes are valid for every conversion pair; unused modes are ignored when absent. :type round: str :param sat: saturation mode for narrow/overflow-prone conversions. - ``True`` (default): Saturate to target range on overflow. - ``False``: No sat; out-of-range values wrap. Appears in: all float->fp8, f32->f16/bf16, all float->int, integer narrowing. Absent from: widening conversions, f16->bf16, s16->f16, s32->f32 (no overflow possible). :type sat: bool :param part: Sub-register half/quarter selector for widening/narrowing conversions. For even/odd 2-way splits: ``"PART_EVEN"`` (0), ``"PART_ODD"`` (1). For fp8/fp4/int4 4-way splits: ``"PART_P0"`` (0), ``"PART_P1"`` (1), ``"PART_P2"`` (2), ``"PART_P3"`` (3). Not needed for same-width conversions (e.g. f16<->bf16, s32->f32, f32->s32, f16->s16). :type part: int :param mode: Write mode for masked-off lanes. - ``"MODE_ZEROING"`` (default): Inactive lanes are set to zero. - ``"MODE_MERGING"``: Inactive lanes preserve their prior value (only on mask-less int->int paths; 920R1 only for mask paths). :type mode: str :returns: * *VReg* -- Destination vector with the converted elements. * *Supported conversion pairs (simplified)* * *----------------------------------------* * **\* Float->Int** (*f32->s64/s32/s16, f16->s32/s16/s8/u8, bf16->s32.*) * **\* Float->Float** (*f32->f16/bf16/fp8, f16->fp8/bf16/f32, bf16->fp8/f16/f32/fp4.*) * **\* Int->Float** (*s16/s32/s64->f16/f32, s8/u8->f16.*) * **\* Int->Int** (*Most s/u{8,16,32,64} widening/narrowing pairs.*) .. py:function:: vsel(src0, src1, mask) Bitwise select: mask ? src0 : src1. .. py:function:: vselr(src, index) Select lanes from src using per-lane indices. .. py:function:: vmaxs(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: vmins(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: vmuls(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: vadds(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: vshls(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: vshrs(src, scalar, mask=None, mode='MODE_ZEROING') .. py:function:: mem_bar(mem_type) Memory barrier: mem_bar(mem_type), e.g. mem_bar(VST_VLD).