@@ -340,7 +340,7 @@ public:
|
||||
base_args.epilogue.thread,
|
||||
reinterpret_cast<const ElementC*>(tensor_c_iter.data()),
|
||||
tensor_c_iter.stride(),
|
||||
reinterpret_cast<const ElementD*>(tensor_d_iter.data()),
|
||||
reinterpret_cast<ElementD*>(tensor_d_iter.data()),
|
||||
tensor_d_iter.stride()
|
||||
};
|
||||
|
||||
|
||||
@@ -82,7 +82,7 @@ struct DistributedGemmKernelWrapper<
|
||||
using BaseArguments = typename BaseKernel::Arguments;
|
||||
using BaseParams = typename BaseKernel::Params;
|
||||
|
||||
static_assert(BaseKernel::ArchTag::kMinComputeCapability == 90, "DistGEMM only supports Hopper GEMMs for now.");
|
||||
//static_assert(BaseKernel::ArchTag::kMinComputeCapability == 90, "DistGEMM only supports Hopper GEMMs for now.");
|
||||
static_assert(not cute::is_same_v<typename BaseKernel::ElementC, void>, "DistributedGEMM epilogues must have a source.");
|
||||
|
||||
using ElementFlag = uint32_t;
|
||||
|
||||
Reference in New Issue
Block a user