CUTLASS 3.3.0 (#1167)
* Release 3.3.0 Adds support for mixed precision GEMMs On Hopper and Ampere Adds support for < 16B aligned GEMMs on Hopper Enhancements to EVT Enhancements to Python interface Enhancements to Sub-byte type handling in CuTe Several other bug-fixes and performance improvements. * minor doc update
This commit is contained in:
@@ -37,7 +37,7 @@ import subprocess
|
||||
|
||||
import torch
|
||||
|
||||
from cutlass import (
|
||||
from cutlass_library import (
|
||||
DataType,
|
||||
DataTypeSize,
|
||||
GemmUniversalMode,
|
||||
@@ -49,7 +49,6 @@ from cutlass import (
|
||||
|
||||
from cutlass.backend import compiler
|
||||
from cutlass.backend.gemm_operation import GemmArguments, GemmOperationUniversal
|
||||
from cutlass.backend.memory_manager import get_allocated_size
|
||||
from cutlass.backend.reduction_operation import ReductionArguments, ReductionOperation
|
||||
from cutlass.shape import GemmCoord, MatrixCoord
|
||||
from cutlass.utils.datatypes import torch_type
|
||||
@@ -65,16 +64,6 @@ class GemmUniversalLauncher:
|
||||
compiler_mode= "nvcc",
|
||||
**kwargs,
|
||||
) -> None:
|
||||
# Create the reduction kernel, if needed
|
||||
self.reduction_operation: ReductionOperation = ReductionOperation(
|
||||
shape=MatrixCoord(4, 32 * operation.C.alignment),
|
||||
C=operation.C,
|
||||
element_accumulator=operation.tile_description.math_instruction.element_accumulator,
|
||||
element_compute=operation.epilogue_functor.element_epilogue,
|
||||
epilogue_functor=operation.epilogue_functor,
|
||||
count=operation.C.alignment,
|
||||
)
|
||||
|
||||
self.math_operation = operation.tile_description.math_instruction.math_operation
|
||||
self.verification = verification
|
||||
|
||||
@@ -88,19 +77,26 @@ class GemmUniversalLauncher:
|
||||
op_list = [operation]
|
||||
if operation.arch < 90:
|
||||
# Split K via Python is currently only supported for pre-SM90 kernels
|
||||
self.reduction_operation: ReductionOperation = ReductionOperation(
|
||||
shape=MatrixCoord(4, 32 * operation.C.alignment),
|
||||
C=operation.C,
|
||||
element_accumulator=operation.tile_description.math_instruction.element_accumulator,
|
||||
element_compute=operation.epilogue_functor.element_epilogue,
|
||||
epilogue_functor=operation.epilogue_functor,
|
||||
count=operation.C.alignment,
|
||||
)
|
||||
op_list.append(self.reduction_operation)
|
||||
|
||||
compiler.add_module(op_list, bypass_cache=False)
|
||||
|
||||
self.operation = operation
|
||||
|
||||
self.dtype_A = torch_type(operation.A.element)
|
||||
self.dtype_B = torch_type(operation.B.element)
|
||||
self.dtype_A = torch_type(operation.A.element if not self.operation.switched else self.operation.B.element)
|
||||
self.dtype_B = torch_type(operation.B.element if not self.operation.switched else self.operation.A.element)
|
||||
self.dtype_C = torch_type(operation.C.element)
|
||||
self.dtype_D = torch_type(operation.C.element)
|
||||
self.dtype_D = torch_type(operation.epilogue_functor.element_output)
|
||||
|
||||
accumulator_size = DataTypeSize[operation.tile_description.math_instruction.element_accumulator]
|
||||
element_size = DataTypeSize[operation.A.element]
|
||||
element_size = min(DataTypeSize[operation.A.element], DataTypeSize[operation.B.element])
|
||||
|
||||
if element_size == 1:
|
||||
self.rand_max = 1
|
||||
@@ -154,7 +150,18 @@ class GemmUniversalLauncher:
|
||||
def reference(self, problem_size, tensor_A, tensor_B, tensor_C, alpha, beta):
|
||||
# If any tensor is on CPU, place all tensors on CPU unless only
|
||||
# tensor C is on CPU
|
||||
devices = [x.device.type for x in [tensor_A, tensor_B, tensor_C]]
|
||||
# Handle mixed-input cases by casting to the larger data type and overriding
|
||||
# to whatever the data type of the larger type is
|
||||
if self.dtype_A != self.dtype_B:
|
||||
if DataTypeSize[self.operation.A.element] < DataTypeSize[self.operation.B.element]:
|
||||
tensor_A = tensor_A.to(self.dtype_B).to(tensor_B.device)
|
||||
else:
|
||||
tensor_B = tensor_B.to(self.dtype_A).to(tensor_A.device)
|
||||
|
||||
devices = [x.device.type for x in [tensor_A, tensor_B]]
|
||||
if tensor_C is not None:
|
||||
devices.append(tensor_C.device.type)
|
||||
|
||||
if "cpu" in devices and devices != ["cuda", "cuda", "cpu"]:
|
||||
device = torch.device("cpu")
|
||||
else:
|
||||
@@ -162,14 +169,17 @@ class GemmUniversalLauncher:
|
||||
|
||||
tensor_A = tensor_A.to(device)
|
||||
tensor_B = tensor_B.to(device)
|
||||
tensor_C = tensor_C.to(device)
|
||||
if tensor_C is not None:
|
||||
tensor_C = tensor_C.to(device)
|
||||
|
||||
dtype = torch_type(self.compute_type)
|
||||
alpha_torch = torch.tensor([alpha], device=device).to(dtype)
|
||||
beta_torch = torch.tensor([beta], device=device).to(dtype)
|
||||
|
||||
tmp = tensor_A @ tensor_B
|
||||
tensor_D_ref = (alpha_torch * tmp) + (tensor_C * beta_torch)
|
||||
tensor_D_ref = (alpha_torch * tmp)
|
||||
if tensor_C is not None:
|
||||
tensor_D_ref += (tensor_C * beta_torch)
|
||||
return tensor_D_ref.to(self.dtype_D)
|
||||
|
||||
def run(self, mode, problem_size, batch_count=1, split_k_slices=1, alpha=1.0, beta=0.0):
|
||||
@@ -199,12 +209,22 @@ class GemmUniversalLauncher:
|
||||
self.dtype_B,
|
||||
self.operation.B.layout if not self.operation.switched else transpose(self.operation.A.layout),
|
||||
)
|
||||
tensor_C, tensor_C_ref = self.uniform_init(
|
||||
if self.dtype_C is not None:
|
||||
tensor_C, tensor_C_ref = self.uniform_init(
|
||||
(true_batch_count, problem_size.m, problem_size.n),
|
||||
self.dtype_C,
|
||||
self.operation.C.layout if not self.operation.switched else transpose(self.operation.C.layout),
|
||||
)
|
||||
else:
|
||||
tensor_C = None
|
||||
tensor_C_ref = None
|
||||
|
||||
tensor_D, _ = self.uniform_init(
|
||||
(true_batch_count, problem_size.m, problem_size.n),
|
||||
self.dtype_C,
|
||||
self.dtype_D,
|
||||
self.operation.C.layout if not self.operation.switched else transpose(self.operation.C.layout),
|
||||
)
|
||||
tensor_D = torch.zeros_like(tensor_C)
|
||||
tensor_D = torch.zeros_like(tensor_D)
|
||||
|
||||
if self.compute_type in [DataType.s8, DataType.s32, DataType.u8, DataType.u32]:
|
||||
alpha = int(alpha)
|
||||
@@ -248,6 +268,10 @@ class GemmUniversalLauncher:
|
||||
if self.verification:
|
||||
if mode == GemmUniversalMode.GemmSplitKParallel:
|
||||
reduction_arguments.sync()
|
||||
|
||||
# Free memory allocated by args because we are not
|
||||
# calling `arguments.sync()` in this case (which will free memory)
|
||||
arguments.free()
|
||||
else:
|
||||
arguments.sync()
|
||||
tensor_D_ref = self.reference(
|
||||
@@ -274,9 +298,6 @@ class GemmUniversalLauncher:
|
||||
if mode == GemmUniversalMode.GemmSplitKParallel:
|
||||
del reduction_arguments
|
||||
|
||||
cur_size = get_allocated_size()
|
||||
assert cur_size == 0, f"{cur_size} B of memory were not released after this run"
|
||||
|
||||
return passed
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user