v3.9 update (#2203)
* v3.9 update * voidD --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -85,7 +85,7 @@ void OperationTable::append(Manifest const &manifest) {
|
||||
|
||||
block_scaled_gemm_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
|
||||
|
||||
|
||||
// insert all gemm operation into operation table
|
||||
if (desc.kind == OperationKind::kGemm) {
|
||||
@@ -121,6 +121,68 @@ void OperationTable::append(Manifest const &manifest) {
|
||||
gemm_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
|
||||
// insert all grouped gemm operation into operation table
|
||||
if (desc.kind == OperationKind::kGroupedGemm) {
|
||||
GroupedGemmDescription const &grouped_gemm_desc = static_cast<GroupedGemmDescription const &>(desc);
|
||||
GemmDescription const &gemm_desc = grouped_gemm_desc.gemm;
|
||||
|
||||
int cc = gemm_desc.tile_description.minimum_compute_capability;
|
||||
|
||||
int alignment = std::max(std::max(
|
||||
gemm_desc.A.alignment, gemm_desc.B.alignment), gemm_desc.C.alignment);
|
||||
|
||||
GemmPreferenceKey preference_key(cc, alignment);
|
||||
|
||||
Operation const *op = operation.get();
|
||||
|
||||
if (!grouped_gemm_desc.block_scales.has_value()) {
|
||||
GemmFunctionalKey functional_key(
|
||||
gemm_desc.provider,
|
||||
gemm_desc.gemm_kind,
|
||||
gemm_desc.tile_description.math_instruction.element_accumulator,
|
||||
gemm_desc.element_epilogue,
|
||||
gemm_desc.A.element,
|
||||
gemm_desc.A.layout,
|
||||
gemm_desc.transform_A,
|
||||
gemm_desc.B.element,
|
||||
gemm_desc.B.layout,
|
||||
gemm_desc.transform_B,
|
||||
gemm_desc.C.element,
|
||||
gemm_desc.C.layout,
|
||||
gemm_desc.D.element,
|
||||
gemm_desc.D.layout
|
||||
);
|
||||
|
||||
gemm_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
else {
|
||||
const BlockScaleDescription &block_scale_desc = grouped_gemm_desc.block_scales.value();
|
||||
BlockScaledGemmFunctionalKey functional_key(
|
||||
gemm_desc.provider,
|
||||
gemm_desc.gemm_kind,
|
||||
gemm_desc.kind,
|
||||
gemm_desc.tile_description.math_instruction.element_accumulator,
|
||||
gemm_desc.element_epilogue,
|
||||
gemm_desc.A.element,
|
||||
gemm_desc.A.layout,
|
||||
block_scale_desc.SFA.element,
|
||||
gemm_desc.B.element,
|
||||
gemm_desc.B.layout,
|
||||
block_scale_desc.SFB.element,
|
||||
gemm_desc.C.element,
|
||||
gemm_desc.C.layout,
|
||||
gemm_desc.D.element,
|
||||
gemm_desc.D.layout,
|
||||
block_scale_desc.SFD.element,
|
||||
block_scale_desc.SFD.layout,
|
||||
block_scale_desc.SFVecSize,
|
||||
block_scale_desc.EpilogueSFVecSize
|
||||
);
|
||||
|
||||
block_scaled_gemm_operations[functional_key][preference_key].push_back(op);
|
||||
}
|
||||
}
|
||||
|
||||
// insert all conv2d or conv3d operation into operation table
|
||||
if (desc.kind == OperationKind::kConv2d || desc.kind == OperationKind::kConv3d) {
|
||||
auto &conv_desc = static_cast<library::ConvDescription const &>(desc);
|
||||
|
||||
Reference in New Issue
Block a user