Arguments and Operands#

An operation is expressed as a RuntimeArguments object that describes both the operation kind (e.g. GEMM) and the operands it runs on. Each operation has its own RuntimeArguments subclass. Operands are described by Operand subclasses such as DenseTensor and ScaledOperand. Tensor- and numeric-typed fields within Arguments and Operands accept any object matching the TensorLike / NumericLike protocols (e.g. torch.Tensor, cute.Tensor, cutlass.Numeric, torch.dtype).

RuntimeArguments#

class cutlass.operators.arguments.base.RuntimeArguments(
*,
performance: PerformanceControls | None = None,
)#

Bases: object

Describes the operands and all other arguments passed to an Operator at runtime.

It contains runtime operands (usually tensors) passed to the operation, as well as any custom epilogue fusions, and runtime performance controls.

This is an abstract base class, whose subclass describes the operation type itself (e.g. GemmArguments, GroupedGemmArguments, etc.). All operators that implement the same operation type accept the same RuntimeArguments subclass.

performance: PerformanceControls | None = None#

Optional runtime performance controls passed to the Operator

class cutlass.operators.arguments.base.PerformanceControls#

Bases: object

Optional runtime performance controls passed to the Operator.

Some operators may support performance options that can be controlled at runtime. This class is the general container for all such controls.

class cutlass.operators.arguments.gemm.GemmProblemSize(M: int, N: int, K: int, L: int)#

Bases: NamedTuple

Problem size for a GEMM operation.

A GEMM with problem size (M, N, K, L) has operands with the following shapes (batch, rows, columns):

  • A: (L, M, K)

  • B: (L, K, N)

  • out: (L, M, N)

M: int#

Number of rows in A and out

N: int#

Number of columns in B and out

K: int#

Number of columns in A and rows in B

L: int#

Number of batches of matrix multiplications

class cutlass.operators.arguments.gemm.GemmArguments(
A: TensorLike | Operand,
B: TensorLike | Operand,
out: TensorLike | Operand,
accumulator_type: NumericLike,
epilogue: EpilogueArguments | None = None,
)#

Bases: RuntimeArguments

Arguments for a Generalized Matrix Multiplication (GEMM) operation: out = A @ B.

The tensors must be all rank-3 or all rank-2. * L: Number of batches * M: Number of rows in A and out * K: Number of columns in A and rows in B * N: Number of columns in B and out

For convenience, construction for A, B, and out that are operands to a dense GEMM can be passed in without wrapping them in DenseTensor.

GemmArguments(A, B, out, accumulator_type)
# is equivalent to:
GemmArguments(DenseTensor(A), DenseTensor(B), DenseTensor(out), accumulator_type)

Other operand types must explicitly wrap tensors in a Operand subclass. For example, a scaled GEMM can be constructed as:

GemmArguments(
    ScaledOperand(A, ScaleATensor, scale_mode, scale_swizzle),
    ScaledOperand(B, ScaleBTensor, scale_mode, scale_swizzle),
    out, # No need to wrap in a `DenseTensor`
    accumulator_type,
)
A: Operand#

Input tensor A of shape (L, M, K) or (M, K)

B: Operand#

Input tensor B of shape (L, K, N) or (K, N)

out: Operand#

Output tensor C of shape (L, M, N) or (M, N)

accumulator_type: NumericLike#

Data type of the accumulator

epilogue: EpilogueArguments | None#

Optional custom epilogue fusion to be performed after the GEMM

property problem_size: GemmProblemSize#

Problem size for a GEMM operation.

class cutlass.operators.arguments.grouped_gemm.IndexPtrGroupedGemmProblemSize(
M: int,
N: int,
K: int,
G: int,
offsets: tuple[int, ...],
offsets_along: Literal['m', 'n', 'k'],
)#

Bases: NamedTuple

Problem shape for contiguous-offset grouped GEMM.

Exactly one of M, N, or K is packed across groups, as selected by offsets_along. Each value in offsets is the exclusive endpoint of one group along that dimension, matching the representation accepted by IndexPtrGroupedGemmArguments.

M: int#

M extent, summed across groups when offsets_along == "m".

N: int#

N extent, summed across groups when offsets_along == "n".

K: int#

K extent, summed across groups when offsets_along == "k".

G: int#

Number of independent GEMMs.

offsets: tuple[int, ...]#

Exclusive endpoint for each group along offsets_along.

offsets_along: Literal['m', 'n', 'k']#

Logical GEMM dimension partitioned by offsets.

class cutlass.operators.arguments.grouped_gemm.IndexPtrGroupedGemmArguments(
A: TensorLike | Operand,
B: TensorLike | Operand,
out: TensorLike | Operand,
accumulator_type: NumericLike,
offsets: TensorLike | Operand,
offsets_along: Literal['m', 'n', 'k'],
epilogue: EpilogueArguments | None = None,
)#

Bases: RuntimeArguments

Arguments for a Grouped GEMM with contiguous tensors concatenated along M, N, or K.

A grouped GEMM performs a series of independent GEMM operations, where each GEMM can have different matrix dimensions. In this particular class, we represent grouped GEMMs where exactly one dimension can vary across groups; the other two remain constant. All operands are dense tensors, packed and contiguous. For the operand with the varying dimension, its total size is sum(changing_dim) * fixed_dims.

offsets describes the problem boundaries along offsets_along. It is a tensor delineating the ending positions of each problem in the group.

When offsets_along == "m" the operands and computation are:

A       : DenseTensor of shape (TotalM, K)   # TotalM = sum(M) over groups
B       : DenseTensor of shape (group_count, K, N)
out     : DenseTensor of shape (TotalM, N)
offsets : DenseTensor of shape (group_count,) = [M0, M0+M1, ..., TotalM]

start = 0
for x in range(group_count):
    end = offsets[x]
    out[start:end, :] = A[start:end, :] @ B[:, :, x]
    start = end

This is equivalent to Pytorch’s 2Dx3D Grouped GEMM.

Rows of A and out are sliced by offsets; B is per group:

     A (TotalM x K)              B[:,:,x]  (K x N, per group)         out (TotalM x N)
             K                          N                                  N
      <--------->                  <------->                          <------->
  0   ┌───────────┐  ┐            ┌─────────┐ ┐                     ┌─────────┐ ┐
      │    A0     │  │ M0   ─ @ ─ │   B0    │ │ K   ────────►       │  out0   │ │ M0
 M0   ├───────────┤  ┤            └─────────┘ ┘                     ├─────────┤ ┤
      │    A1     │  │ M1   ─ @ ─ ┌─────────┐ ┐                     │  out1   │ │ M1
M0+M1 ├───────────┤  ┤            │   B1    │ │ K   ────────►       ├─────────┤ ┤
      │    A2     │  │ M2   ─ @ ─ └─────────┘ ┘                     │  out2   │ │ M2
TotalM└───────────┘  ┘            ┌─────────┐ ┐                     └─────────┘ ┘
        ▲                         │   B2    │ │ K   ────────►
      offsets slice rows          └─────────┘ ┘
      [0:M0], [M0:M0+M1], ...

  per group g:   out[start:end, :]  =  A[start:end, :]  @  B[:, :, g]
                    (Mg x N)             (Mg x K)           (K x N)

When offsets_along == "n" the operands and computation are:

A       : DenseTensor of shape (group_count, M, K)
B       : DenseTensor of shape (K, TotalN)   # TotalN = sum(N) over groups
out     : DenseTensor of shape (M, TotalN)
offsets : DenseTensor of shape (group_count,) = [N0, N0+N1, ..., TotalN]

start = 0
for x in range(group_count):
    end = offsets[x]
    out[:, start:end] = A[:, :, x] @ B[:, start:end]
    start = end

This is equivalent to Pytorch’s 3Dx2D Grouped GEMM.

Columns of B and out are sliced by offsets; A is per group:

A[:,:,x] (M x K, per group)       B (K x TotalN)             out (M x TotalN)
        K                          N0    N1    N2             N0    N1    N2
   <─────────>                  <────><────><────>        <────><────><────>
┌─────────────┐ ┐              ┌─────┬─────┬─────┐        ┌─────┬─────┬─────┐
│     A0      │ │ M  ─ @ ─►    │     │     │     │        │     │     │     │
└─────────────┘ ┘            K │ B0  │ B1  │ B2  │ K    M │out0 │out1 │out2 │ M
┌─────────────┐ ┐              │     │     │     │        │     │     │     │
│     A1      │ │ M  ─ @ ─►    └─────┴─────┴─────┘        └─────┴─────┴─────┘
└─────────────┘ ┘                 ▲     ▲     ▲
┌─────────────┐ ┐              offsets slice columns
│     A2      │ │ M  ─ @ ─►     [0:N0], [N0:N0+N1], ...
└─────────────┘ ┘

per group g:   out[:, start:end]  =  A[:, :, g]  @  B[:, start:end]
                  (M x Ng)             (M x K)        (K x Ng)

When offsets_along == "k" each K-slice is an independent GEMM whose M x N product is written to its own section in the output tensor:

A       : DenseTensor of shape (M, TotalK)          # TotalK = sum(K) over groups
B       : DenseTensor of shape (TotalK, N)
out     : DenseTensor of shape (group_count, M, N)  # one M x N result per group
offsets : DenseTensor of shape (group_count,) = [K0, K0+K1, ..., TotalK]

start = 0
for x in range(group_count):
    end = offsets[x]
    out[:, :, x] = A[:, start:end] @ B[start:end, :]
    start = end

This is equivalent to Pytorch’s 2Dx2D Grouped GEMM.

offsets slice columns of A and rows of B; each slice is one group:

  A (M x TotalK)                        B (TotalK x N)           out (group_count, M, N)
       K0     K1    K2                       N                        N
  <────><─────><─────>                     <──────>                <──────>
  ┌─────┬──────┬──────┐  ┐          0     ┌─────────┐ ┐     x=0  ┌─────────┐ ┐
M │ A0  │  A1  │  A2  │  │ M              │   B0    │ │ K0       │  A0@B0  │ │ M
  └─────┴──────┴──────┘  ┘         K0     ├─────────┤ ┤          └─────────┘ ┘
     │      │      │                      │   B1    │ │ K1  x=1  ┌─────────┐ ┐
     │      │      └────@───► B2    K0+K1 ├─────────┤ ┤          │  A1@B1  │ │ M
     │      └───────────@───► B1          │   B2    │ │ K2       └─────────┘ ┘
     └──────────────────@───► B0   TotalK └─────────┘ ┘     x=2  ┌─────────┐ ┐
                                                                 │  A2@B2  │ │ M
                                                                 └─────────┘ ┘

  per group x:  out[:, :, x] = A[:, start:end] @ B[start:end, :]
                  (M x N)         (M x Kx)         (Kx x N)
A: DenseTensor#

Dense input tensor A.

The shape of A is (sum(M across all groups or TotalM), K) if offsets_along is “m” or (M, sum(K across all groups or TotalK)) if offsets_along is “k” or (group_count, M, N) if offsets_along is “n” where M, N, K are the dimensions of the GEMM problem.

B: DenseTensor#

Dense input tensor B.

The shape of B is (K, sum(N across all groups or TotalN)) if offsets_along is “n” or (sum(K across all groups or TotalK), N) if offsets_along is “k” or (group_count, K, N) if offsets_along is “m” where M, N, K are the dimensions of the GEMM problem.

out: DenseTensor#

Dense Output tensor.

The shape of out is (sum(M across all groups or TotalM), N) if offsets_along is “m” or (M, sum(N across all groups or TotalN)) if offsets_along is “n” or (group_count, M, N) if offsets_along is “k” where M, N, K are the dimensions of the GEMM problem.

accumulator_type: NumericLike#

Data type of the accumulator.

offsets: DenseTensor#

Integer tensor describing grouped-problem boundaries.

offsets_along: Literal['m', 'n', 'k']#

Logical GEMM dimension partitioned by offsets.

epilogue: EpilogueArguments | None#

Optional custom epilogue fusion to perform after the GEMM.

class cutlass.operators.arguments.grouped_gemm.GroupedGemmArguments(
A: TensorLike | Operand,
B: TensorLike | Operand,
out: TensorLike | Operand,
accumulator_type: NumericLike,
offsets: TensorLike | Operand,
epilogue: EpilogueArguments | None = None,
)#

Bases: IndexPtrGroupedGemmArguments

Deprecated compatibility interface for M-offset grouped GEMM.

Use IndexPtrGroupedGemmArguments with offsets_along="m".

class cutlass.operators.arguments.epilogue.EpilogueArguments(epilogue_fn: Callable | str, **kwargs)#

Bases: object

Describes a user-defined epilogue that is fused on top of the operation described by the primary RuntimeArguments.

An epilogue fusion is a custom function that performs tensor-level transformations on the result of a matrix multiplication. This transformation is fused into the kernel’s epilogue, which stores the final output.

EpilogueArguments encapsulates the epilogue function describing the transformation along with its arguments.

To support flexible definition of epilogues, EpilogueArguments is defined generically as taking in an epilogue_fn and additional kwargs.

Under the hood, the AST for epilogue_fn is parsed to determine the operands and outputs of the epilogue. kwargs must contain Tensors or scalars for all operands and outputs in the provided epilogue.

Structure of ``epilogue_fn``

The epilogue_fn is a function describing the custom transformation on the accumulator tensor (intermediate result before the epilogue).

The general structure of these functions is:

def custom_epi_name(accum, *args) -> TensorType | tuple[TensorType, ...]:
    '''Compute the epilogue.

    # Args:
        accum (TensorType): Result of the primary operation (e.g.
                        ``A @ B`` for a GEMM) before the epilogue.
        *args: Additional tensors or scalars used in the epilogue
                (e.g. aux tensors).

    # Returns:
        At least one tensor resulting from the epilogue computation.
    '''
    # Do some compute
    return D  # and potentially other values

epilogue_fn must be a Python callable (or its string representation) that must satisfy the following constraints:

  • Takes a first positional argument named accum – the result of the operation just before the epilogue. For a GEMM, accum = A @ B.

  • Returns at least one tensor resulting from the epilogue. Currently the return list must contain at least one output named D.

  • Each argument following accum is a tensor or scalar to be loaded.

  • Each variable in the return statement is a tensor or scalar to be stored.

  • Operations are represented in static single assignment (SSA) form. This means that each variable can be assigned exactly once.

The epilogue is parsed, never executed

The function body is a small domain-specific language: its source is parsed (a callable’s source is read back with inspect.getsource, so callables and source strings are equivalent) and names are matched against the fixed vocabulary below. They do not resolve to Python objects – max here is not builtins.max – and calling arbitrary Python functions is not supported.

  • Arithmetic: +, -, *, / – elementwise, with scalar and row/column-vector broadcasting of the non-accum operand.

  • Elementwise functions: relu, tanh, sigmoid, silu, hardswish, gelu, exp, maximum, minimum, multiply_add, identity.

  • Layout: permute, reshape.

  • Reductions: sum, max and min (see below).

Reductions

A returned value may reduce the accumulator instead of storing a tile; the reduction is fused into the same kernel, so the M x N result is never re-read:

def epi_with_reductions(accum):
    D = accum
    total = sum(accum)              # scalar: all axes folded
    col_max = max(accum, dim=[0])   # per-column: destination (N,)
    row_max = max(accum, dim=[1])   # per-row: column-major (M, 1)
    return D, total, col_max, row_max

The folded-vs-kept axes are derived from the layout of the destination tensor passed in kwargs (the dim= keyword documents intent): a single-element tensor selects the scalar reduction, a 1-D (N,) tensor keeps the columns, and a column-major (M, 1) tensor keeps the rows – a contiguous (M, 1) or a bare (M,) is ambiguous and rejected. Reductions combine onto the destination with atomics, so initialise it to the operation’s identity (0 for sum, -inf for max, +inf for min), or to an existing partial result to accumulate into it. Batch axes are folded into the same destination; per-batch reduction outputs are not supported yet.

Choosing data transports

By default the kernel decides how each operand is moved between global memory and the compute; wrap a tensor kwarg in Load or Store (e.g. C=Load(C, via=Transport.ASYNC_GMEM_LOAD)) to select a specific transport. See those classes below.

Structure of ``kwargs``

kwargs must contain sample Tensors or scalars for all operands and outputs in the provided epilogue. For example, with an epilogue of:

def my_epi(accum, alpha, C, beta):
    F = (accum * alpha) + (C * beta)
    D = relu(F)
    return D, F

A user would need to construct epilogue arguments as follows:

epi_args = EpilogueArguments(
    my_epi,
    alpha=..., C=..., beta=..., D=..., F=...
)
copy() → EpilogueArguments#

Return a copy of these arguments. Does not copy the underlying tensors.

Constructing a RuntimeArguments traces the epilogue and converts its tensors, both of which mutate the object. Operating on a copy keeps a caller-supplied EpilogueArguments reusable across several operations, mirroring how GemmArguments copies its A/B/out operands.

Returns:

An independent, untraced copy.

Return type:

EpilogueArguments

property parameters: list[cute.Tensor | Numeric]#

Returns the list of input and output parameters of the epilogue.

property parameter_names: list[str]#

Returns the list of names of the input and output parameters of the epilogue.

to_tensor_wrappers(permute: list[int] | None = None)#

Converts the input and output parameters of the epilogue to TensorWrappers.

trace(
accumulator_shape: tuple[int, ...],
accumulator_type: Numeric,
)#

Traces the epilogue function and generates an internal representation of the epilogue.

Parameters:
  • accumulator_shape (tuple[int, ...]) – The shape of the accumulator tensor. For example, for a GEMM, this would be the shape of the output tensor.

  • accumulator_type (Numeric) – The datatype of the accumulator tensor.

class cutlass.operators.arguments.epilogue.Transport(value)#

Bases: Enum

How a supplemental epilogue tensor moves between GMEM and registers.

TMA stages through shared memory; the *_GMEM_* variants use direct GMEM addressing, either synchronous to the issuing thread (SYNC_GMEM_LOAD / SYNC_GMEM_STORE) or asynchronous through shared memory via cp.async (ASYNC_GMEM_LOAD).

class cutlass.operators.arguments.epilogue.Load(
tensor: Any,
via: Transport | str = Transport.TMA,
num_bits_per_copy: int | None = None,
)#

Bases: object

Describe how an epilogue input tensor is read.

Wrap an EpilogueArguments tensor kwarg to override the default (TMA) read transport, e.g. C=ops.Load(C, via=ops.Transport.ASYNC_GMEM_LOAD). Passing a bare tensor is equivalent to Load(tensor) (a TMA read).

Parameters:
  • tensor (TensorLike) – The tensor to read.

  • via (Transport | str) – Read transport; one of TMA, SYNC_GMEM_LOAD or ASYNC_GMEM_LOAD. Strings are accepted.

  • num_bits_per_copy (int | None) – Transaction width for non-TMA transports; None auto-derives it.

class cutlass.operators.arguments.epilogue.Store(
tensor: Any,
via: Transport | str = Transport.TMA,
num_bits_per_copy: int | None = None,
)#

Bases: object

Describe how an epilogue output tensor is written.

Wrap an EpilogueArguments tensor kwarg to override the default (TMA) write transport, e.g. D=ops.Store(D, via=ops.Transport.SYNC_GMEM_STORE). Passing a bare tensor is equivalent to Store(tensor) (a TMA write).

Parameters:
  • tensor (TensorLike) – The tensor to write.

  • via (Transport | str) – Write transport; one of TMA or SYNC_GMEM_STORE. Strings are accepted.

  • num_bits_per_copy (int | None) – Transaction width for the direct store; None auto-derives it.

Operands#

class cutlass.operators.arguments.base.Operand#

Bases: ABC

Base class for all operands to Operators, which encapsulates one or more TensorLike objects.

In the most basic case, an Operand enacapsulates a single tensor.

In more complex cases, an Operand may encapsulate multiple tensors that encapsulate a single logical operand. For instance, a ScaledOperand encapsulates a quantized and scale tensor, that together reconstruct the logical value of the operand.

final copy() → Operand#

Returns a copy of the operand. Does not copy the underlying tensor.

class cutlass.operators.arguments.operand.DenseTensor(tensor: cutlass.operators.typing.TensorLike)#

Bases: Operand

An operand encapsulating a simple, single dense tensor.

class cutlass.operators.arguments.operand.ScaledOperand(
quantized: TensorLike | DenseTensor,
scale: TensorLike | DenseTensor,
mode: ScaleMode | tuple[int, ...],
swizzle: ScaleSwizzleMode,
)#

Bases: Operand

An operand whose logical value is a quantized tensor multiplied by a tensor of scale factors.

A ScaledOperand encapsulates physical quantized and scale tensors that together represent the logical value scale * quantized: each entry of scale broadcasts and multiplies a contiguous block of quantized, where the block shape is given by mode.

This Operand is typically used to express narrow-precision formats (such as OCP MXFP8 / MXFP4 and NVIDIA NVFP4), which store the data as a quantized narrow-precision tensor multiplied by a separate tensor of scale factors that recover dynamic range.

The scale mode and swizzle describe the physical layout of the scale tensor.

The mode describes how scale factors are broadcast over the quantized tensor, and therefore determines its size. It’s usually dictated by the data format you seek to represent – MXFP8 / MXFP4 / NVFP4 require specific scale modes. See ScaleMode for more details.

The swizzle describes the in-memory layout of the scale tensor. It’s usually dictated by the hardware architecture – for example, Blackwell block-scaled MMAs require a specific swizzle layout, and most kernels using them require the input to be pre-arranged in this layout. See ScaleSwizzleMode for more details.

Example

quantized_A = torch.randn(M, K, dtype=torch.float8_e4m3fn, device="cuda")
scale_A = torch.randn(M, K // 32, dtype=torch.float8_e8m0fnu, device="cuda")
A = ScaledOperand(
    quantized_A, scale_A, ScaleMode.Blockwise1x32, ScaleSwizzleMode.Swizzle32x4x4,
)

See also

  • ScaleMode: granularity of the scaling, and the format-to-mode mapping.

  • ScaleSwizzleMode: in-memory layout of the scale tensor.

References

quantized: DenseTensor#

Narrow-precision tensor holding the quantized values.

The dequantized logical value is scale * quantized: each scale factor broadcasts over a block of this tensor (block shape given by mode) and multiplies its values.

scale: DenseTensor#

Tensor of scale factors.

Each entry of this tensor broadcasts over a block of quantized and scales it. Its shape and layout are determined together by mode and swizzle.

This tensor must be contiguous and have exactly ScaledOperand.numel_scale() (quantized_shape, mode, swizzle) elements. It must already be arranged in the layout named by swizzle. Its shape is otherwise unconstrained and not validated.

mode: ScaleMode | tuple[int, ...]#

Block shape over which each scale factor is broadcast.

Typically, a block of quantized-tensor elements share the same scale factor value. The logical value of the scaled operand is the result of broadcasting the scale factor along the block shape described by mode, and then multiplying it element-wise with the quantized tensor.

This may be a ScaleMode enum for commonly used block shapes, or a bare (L, M, K) tuple for custom block shapes. For instance, (1, 1, 32) means that each scale factor covers 32 elements along the K axis, and 1 element along the L and M axes.

swizzle: ScaleSwizzleMode#

In-memory layout of the scale factor tensor.

The scale factor tensor is sometimes required to be stored in a specific swizzled layout, often dictated by the hardware MMA’s contract.

static numel_scale(
quantized_shape: tuple[int, ...],
mode: ScaleMode | tuple[int, ...],
swizzle: ScaleSwizzleMode,
) → int#

Return the number of elements expected in scale for an operand of quantized_shape.

quantized_shape is (L, outer, K) for a 3D operand, where outer is M for an A-side scale and N for a B-side scale. One scale factor covers V elements along K, so the scale tensor has logical shape (L, outer, ceil_div(K, V)).

scale.numel() must equal the value returned by this method. The shape of the scale tensor is otherwise unconstrained.

Returned value depends on swizzle:

where V = ScaleMode.numel(mode) is the number of quantized-tensor elements covered by one scale factor.

Parameters:
  • quantized_shape (tuple[int, ...]) – Logical shape of the operand the scale will be paired with, as (L, outer, K) or (outer, K).

  • mode (ScaleMode | tuple[int, ...]) – Block shape of the scaling.

  • swizzle (ScaleSwizzleMode) – In-memory layout of the scale tensor.

Returns:

Required number of scale-factor elements.

Return type:

int

Raises:

ValueError – If quantized_shape is not rank 2 or 3, or swizzle is not a recognised ScaleSwizzleMode.

class cutlass.operators.arguments.operand.ScaleMode(value)#

Bases: Enum

An enum over commonly used block scaling modes for a ScaledOperand.

Each member’s value is a (batch, row, col) tuple giving the shape of the block of quantized-tensor elements that share a scale factor. Each scale factor is broadcast & multiplied over this block of quantized-tensor elements.

For block shapes that are not enumerated here, bare (batch, row, col) tuples are also accepted as ScaledOperand.mode

See also

Blockwise1x16 = (1, 1, 16)#

Scale over (1, 1, 16) block. Used by NVIDIA NVFP4 (FP8 E4M3 scale dtype).

Blockwise1x32 = (1, 1, 32)#

Scale over (1, 1, 32) block. Used by OCP MXFP8 / MXFP4 (E8M0 scale dtype).

static compare(
mode1: ScaleMode | tuple[int, ...],
mode2: ScaleMode | tuple[int, ...],
) → bool#

Return whether two scale modes describe the same scaling.

Allows mixed comparisons between ScaleMode enums and bare tuples, and tolerates differing tuple lengths as long as the longer tuple has leading 1 s for the extra positions (so that, e.g., (1, 1, 16) == (1, 16)).

Parameters:
  • mode1 (ScaleMode | tuple[int, ...]) – First mode.

  • mode2 (ScaleMode | tuple[int, ...]) – Second mode.

Returns:

True if mode1 and mode2 describe the same scaling.

Return type:

bool

static numel(
scale: ScaleMode | tuple[int, ...],
) → int#

Return the number of quantized-tensor elements scaled by one scale factor.

For a scale mode of block shape (L, M, K), this is the block volume L * M * K – the number of contiguous elements in a quantized tensor that share (and are multiplied by) the same scale factor.

Parameters:

scale (ScaleMode | tuple[int, ...]) – Scale mode whose tuple entries give the block shape.

Returns:

Product of the entries of scale (e.g. numel(Blockwise1x32) == 32).

Return type:

int

class cutlass.operators.arguments.operand.ScaleSwizzleMode(value)#

Bases: Enum

In-memory layout of the scale tensor in a ScaledOperand.

Hardware MMAs that consume block-scaled inputs typically require their scale tensors to be stored in a specific non-trivial layout.

ScaleSwizzleMode indicates that scale factors inside a ScaledOperand are already stored in one of those specific layouts.

Operator API validates that the scale tensor is contiguous and has the element count required by (mode, swizzle), but it does not inspect or rearrange the values. The producer of the scale tensor (typically a quantizer) is responsible for writing values in the named layout.

See also

  • ScaledOperand: the operand that consumes this enum.

  • ScaleMode: granularity of the scaling.

  • PyTorch torch.nn.functional.SwizzleType, which this enum closely resembles.

SwizzleNone = 1#

Scale tensor stored in the natural ordering implied by ScaleMode.

For an (L, M, K) operand at mode (1, 1, V), this is a tensor of L * M * (K // V) scale factors with one value per mode block.

Swizzle32x4x4 = 2#

1D block-scaling layout dictated by Blackwell tcgen05.mma block-scaled MMAs.

This is used for MXFP8 / MXFP4 / NVFP4 GEMMs on Blackwell. Under this layout, each tile contains 128x4 scale factors, and the 128 rows are interleaved in groups of 32, matching the warp-group structure of the Blackwell tensor core.

For a quantized tensor of shape (L, M, K) and a scale vector size of V, the scale tensor is required to have L * round_up(M, 128) * round_up(ceil_div(K, V), 4) elements.

References

Type markers#

Type markers for CUTLASS Operator API field annotations.

cutlass.operators.typing.TensorLike#

Tensor-like inputs accepted by the Operator API.

Includes:
  • Any DLPack-compatible tensor (torch.Tensor, jax.Array, numpy.ndarray, …).

  • cutlass.cute.Tensor (CuTe DSL host tensor; does not implement DLPack but is handled natively by TensorWrapper).

  • cutlass.operators.utils.tensor.TensorWrapper (implements DLPack).

alias of _SupportsDLPack | Tensor | TensorWrapper

class cutlass.operators.typing.NumericLike(*args, **kwargs)#

Bases: Protocol

Type marker for fields that accept numeric-like types.

Fields annotated with NumericLike accept the following:
  • cutlass.Numeric

  • torch.dtype