12
345678910
1112131415161718192021
22readonly func BundleCubeMatrixSelected() => boolean23begin24 return _BundleOperation.valid &&25 _BundleOperation.operation_class == BundleOperation_TileMatrix &&26 _BundleOperation.selector_valid &&27 TileMatrixFunctionAssigned(28 UInt(_BundleOperation.selector[4:0]));29end;30
31readonly func BundleTMATMULDataAttributesLegal() => boolean32begin33 if !_BundleDataAttributesPresent then return TRUE; end;34 return _BundleDataAttributes.data_layout == Zeros{5} &&35 _BundleDataAttributes.pad_value == Zeros{2} &&36 _BundleDataAttributes.comparison_mode == Zeros{3} &&37 !_BundleDataAttributes.canonicalize;38end;39
40readonly func BundleTMATMULMasksAgree() => boolean41begin42 var seen = FALSE;43 var selected = Zeros{4};44 for binding = 0 to PTO_BUNDLE_TILE_BINDING_COUNT - 1 do45 if _BundleTileBindings[[binding]].valid then46 let mask = _BundleTileBindings[[binding]].pe_mask;47 if seen && mask != selected then return FALSE; end;48 selected = mask;49 seen = TRUE;50 end;51 end;52 for binding = 0 to 3 do53 if _BundleSharedBindings[[binding]].valid then54 let mask = _BundleSharedBindings[[binding]].pe_mask;55 if seen && mask != selected then return FALSE; end;56 selected = mask;57 seen = TRUE;58 end;59 end;60 return seen;61end;62
63readonly func BundleTMATMULSharedMasksAreZero() => boolean64begin65 for binding = 0 to 3 do66 if _BundleSharedBindings[[binding]].valid &&67 _BundleSharedBindings[[binding]].pe_mask != Zeros{4} then68 return FALSE;69 end;70 end;71 return TRUE;72end;73
74readonly func BundleMatrixPostProcessSourceCount() => integer {0..3}75begin76 return77 (if _BundleFixedPointAttributes.row_max_en &&78 _BundleFixedPointAttributes.row_max_init79 then 1 else 0) +80 (if BundleFPATRModeUsesVectorParameter(81 _BundleFixedPointAttributes.pre_quant_mode)82 then 1 else 0) +83 (if BundleFPATRReluModeUsesVectorParameter(84 _BundleFixedPointAttributes.relu_mode)85 then 1 else 0);86end;87
88readonly func BundleMatrixDestinationCount() => integer {1..3}89begin90 return 1 +91 (if _BundleFixedPointAttributes.row_max_en then 1 else 0) +92 (if _BundleFixedPointAttributes.group_max_en then 1 else 0);93end;94
95readonly func BundleMatrixDynamicBindingsComplete(96 operation: integer {0..PTO_TILE_OPERATION_COUNT-1},97 function: integer {0..31},98 left_type: TileDataType,99 right_type: TileDataType,100 shared_count: integer {0..4}) => boolean101begin102 if !_BundleFixedPointAttributes.valid ||103 !TileMatrixSharedSourceCountLegal(104 function, left_type, right_type, shared_count) then105 return FALSE;106 end;107 let mathematical_sources = TileMatrixLocalMathematicalSourceCount(108 function, left_type, right_type, shared_count) +109 (if _BundleFixedPointAttributes.c_scale_en then 1 else 0);110 let expected_sources = mathematical_sources +111 BundleMatrixPostProcessSourceCount();112 return BundleLocalTileSourceCount() == expected_sources &&113 BundleLocalTileDestinationCount() ==114 BundleMatrixDestinationCount() &&115 BundleTileBindingStreamTerminated() &&116 BundleOperationScalarBindingSchemaLegal(operation);117end;118
119pure func BundleTMATMULCooperativeSelected(120 function: integer {0..31},121 shared_count: integer {0..4}) => boolean122begin123 return shared_count > 0 && !TileMatrixFunctionIsGEMV(function);124end;125
126pure func BundleTMATMULCooperativeMaskValueLegal(mask: bits(4)) => boolean127begin128 return mask == '1111';129end;130
131readonly func BundleTMATMULCooperativeMasksLegal(132 function: integer {0..31},133 shared_count: integer {0..4}) => boolean134begin135 if !BundleTMATMULCooperativeSelected(function, shared_count) then136 return TRUE;137 end;138 for binding = 0 to PTO_BUNDLE_TILE_BINDING_COUNT - 1 do139 if _BundleTileBindings[[binding]].valid &&140 !BundleTMATMULCooperativeMaskValueLegal(141 _BundleTileBindings[[binding]].pe_mask) then142 return FALSE;143 end;144 end;145 for binding = 0 to 3 do146 if _BundleSharedBindings[[binding]].valid &&147 !BundleTMATMULCooperativeMaskValueLegal(148 _BundleSharedBindings[[binding]].pe_mask) then149 return FALSE;150 end;151 end;152 return TRUE;153end;154
155readonly func BundleTMATMULSelectedMask() => bits(4)156begin157 for binding = 0 to PTO_BUNDLE_TILE_BINDING_COUNT - 1 do158 if _BundleTileBindings[[binding]].valid then159 return _BundleTileBindings[[binding]].pe_mask;160 end;161 end;162 return Zeros{4};163end;164
165readonly func BundleTMATMULCurrentPEInactive() => boolean166begin167 if !BundleCubeMatrixSelected() then return FALSE; end;168 let function = UInt(_BundleOperation.selector[4:0]);169 let shared_count = BundleSharedBindingCount();170 if !BundleTMATMULCooperativeSelected(function, shared_count) then171 return FALSE;172 end;173 let group_m = BundleCubeDimensionValue(BundleDimension_LB0);174 if group_m == 0 || group_m > 128 then return FALSE; end;175 return BundleMatrixCooperativeValidM(176 group_m as integer {1..65535}, _CurrentMemoryAgent) == 0;177end;178
179readonly func BundleMatrixPrimaryDestinationCapacityBytes()180 => integer {0,128,256,512,1024,2048,4096,8192,16384,32768,65536,181 131072,262144}182begin183 for binding = 0 to PTO_BUNDLE_TILE_BINDING_COUNT - 1 do184 if _BundleTileBindings[[binding]].valid &&185 _BundleTileBindings[[binding]].destination_valid then186 return BundleTileDestinationSizeBytes(187 binding as BundleTileBindingIndex);188 end;189 end;190 return 0;191end;192
193readonly func BundleMatrixAccumulatorDestinationIndicesDistinct(194 function: integer {0..31}) => boolean195begin196 if !TileMatrixFunctionUsesAccumulator(function) then return TRUE; end;197 let (destination_seen, destination_hand) =198 BundleMatrixPrimaryDestinationHand();199 if !destination_seen then return FALSE; end;200 var accumulator: TileIndex = 0;201 var found = FALSE;202 for binding = 0 to PTO_BUNDLE_TILE_BINDING_COUNT - 1 do203 if !found && _BundleTileBindings[[binding]].valid then204 if _BundleTileBindings[[binding]].source0_valid then205 accumulator = BundleTileArchitecturalSourceIndex(206 binding as BundleTileBindingIndex, FALSE);207 found = TRUE;208 elsif _BundleTileBindings[[binding]].source1_valid then209 accumulator = BundleTileArchitecturalSourceIndex(210 binding as BundleTileBindingIndex, TRUE);211 found = TRUE;212 end;213 end;214 end;215 if !found then return FALSE; end;216 return accumulator != destination_hand;217end;218
219func ExecuteBundleTMATMULOperation() => boolean220begin221 222 223 224 if SelectedBundleTileMaskIsZero() &&225 BundleTMATMULSharedMasksAreZero() then226 return TRUE;227 end;228
229 if !_BundleFixedPointAttributes.valid then230 SetFault(Fault_BundleControl, ReadTPC());231 return FALSE;232 end;233
234 let decoded = DecodeTileOperation(TileDecode_CUBE,235 BundleOperationDecodeCode(_BundleOperation));236 if decoded == PTO_TILE_OPERATION_COUNT then237 SetFault(Fault_IllegalInstruction, ReadTPC());238 return FALSE;239 end;240 let operation = decoded as integer {0..PTO_TILE_OPERATION_COUNT-1};241 let function = UInt(_BundleOperation.selector[4:0]);242 let left_type = TileDataTypeFromEncoding(243 CurrentBundleTileOperationDataTypeCode() as TileDataTypeEncoding);244 let right_type = if _BundleDataAttributesPresent then245 TileDataTypeFromEncoding(246 _BundleDataAttributes.data_type as TileDataTypeEncoding)247 else left_type;248 let shared_count = BundleSharedBindingCount();249 let matrix_types_legal = if TileMatrixFunctionUsesMX(function) then250 TileMXOperandPairLegal(left_type, right_type)251 else252 TileOrdinaryMatrixInputTypesSameClass(left_type, right_type);253 if !matrix_types_legal ||254 !BundleMatrixDynamicBindingsComplete(255 operation, function, left_type, right_type, shared_count) ||256 !BundleTMATMULDataAttributesLegal() ||257 !BundleTMATMULDimensionsLegal(shared_count) ||258 !SelectedBundleTileMasksLegal() ||259 !BundleTMATMULMasksAgree() ||260 !BundleTMATMULCooperativeMasksLegal(function, shared_count) then261 SetFault(Fault_TileLegality, ReadTPC());262 return FALSE;263 end;264
265 let m_raw = BundleCubeDimensionValue(BundleDimension_LB0);266 let n_raw = BundleCubeDimensionValue(BundleDimension_LB1);267 let k_raw = BundleCubeDimensionValue(BundleDimension_LB2);268 let m = m_raw as integer {1..65535};269 let n = n_raw as integer {1..65535};270 let k = k_raw as integer {1..65535};271 if TileMatrixFunctionIsGEMV(function) && m != 1 then272 SetFault(Fault_TileLegality, ReadTPC());273 return FALSE;274 end;275 if !BundleMatrixSharedSchemasLegal(276 function, left_type, right_type,277 m, n, k, shared_count) then278 SetFault(Fault_TileLegality, ReadTPC());279 return FALSE;280 end;281
282 let cooperative = BundleTMATMULCooperativeSelected(283 function, shared_count);284 let valid_m = if cooperative then285 BundleMatrixCooperativeValidM(m, _CurrentMemoryAgent)286 else m;287 if cooperative && valid_m == 0 then288 ConsumeBundleSharedBindings(shared_count as integer {1..4});289 FinalizeBundleTileAttempt(TileExecution_Executed);290 return TRUE;291 end;292 let pe_m = valid_m as integer {1..65535};293 294 295 296 if !PrepareSelectedBundleStage2() then return FALSE; end;297 if !ReuseBundleLocalGenerationDestination() then return FALSE; end;298 if !BundleOperationGPRBindingValuesLegal(operation) then299 SetFault(Fault_TileLegality, ReadTPC());300 return FALSE;301 end;302
303 let mathematical_sources = TileMatrixLocalMathematicalSourceCount(304 function, left_type, right_type, shared_count) +305 (if _BundleFixedPointAttributes.c_scale_en then 1 else 0);306 let result_type = if TileMatrixFunctionUsesMX(function) then307 TileDataType_FP32308 else309 TileOrdinaryMatrixAccumulatorType(left_type, right_type);310 if _BundleFixedPointAttributes.c_scale_en &&311 (!TileMatrixFunctionAllowsCScale(function) ||312 result_type != TileDataType_FP32) then313 SetFault(Fault_TileLegality, ReadTPC());314 return FALSE;315 end;316 if !BundleMatrixPostProcessSourcesLegal(317 mathematical_sources, pe_m, n, result_type) then318 SetFault(Fault_TileLegality, ReadTPC());319 return FALSE;320 end;321 if !BundleMatrixLocalMathematicalSourcesLegal(322 function, left_type, right_type, pe_m, n, k, shared_count,323 result_type, BundleMatrixPrimaryDestinationCapacityBytes()) then324 SetFault(Fault_TileLegality, ReadTPC());325 return FALSE;326 end;327 if !BundleMatrixAccumulatorDestinationIndicesDistinct(function) then328 SetFault(Fault_TileLegality, ReadTPC());329 return FALSE;330 end;331 if _BundleFixedPointAttributes.c_scale_en &&332 !BundleMatrixCScaleDestinationIndicesDistinct(333 (mathematical_sources - 1) as integer {0..8}) then334 SetFault(Fault_TileLegality, ReadTPC());335 return FALSE;336 end;337 let (layout_found, primary_layout) =338 BundleMatrixCooperativeMLayout(339 function, right_type, pe_m, shared_count);340 if !layout_found then341 SetFault(Fault_TileLegality, ReadTPC());342 return FALSE;343 end;344 345 346 347 let allocation_mask = if cooperative then348 BundleMatrixCooperativeCurrentPEMask(m, _CurrentMemoryAgent)349 else BundleTMATMULSelectedMask();350 if !ResolveBundleTMATMULDestination(351 pe_m, n, result_type, TRUE, primary_layout,352 allocation_mask) then353 return FALSE;354 end;355
356 let operands = BundleTileInstructionOperands(operation);357 var left = _Tiles[[0]];358 var right = _Tiles[[0]];359 var left_scale = _Tiles[[0]];360 var right_scale = _Tiles[[0]];361 let left_scale_present = TileMatrixFunctionUsesMX(function) &&362 TileMXInputTypeNeedsScale(left_type);363 let right_scale_present = TileMatrixFunctionUsesMX(function) &&364 TileMXInputTypeNeedsScale(right_type);365 var accumulator: TileIndex = operands.destination0;366 var bias: TileIndex = operands.destination0;367 var c_scale: TileIndex = operands.destination0;368 var local_ordinal: integer {0..6} = 0;369 var shared_ordinal: integer {0..4} = 0;370
371 if TileMatrixFunctionUsesAccumulator(function) then372 accumulator = BundleMatrixSourceAt(373 local_ordinal as integer {0..8});374 local_ordinal = (local_ordinal + 1) as integer {0..6};375 end;376
377 if shared_count == 0 then378 left = _Tiles[[BundleMatrixSourceAt(379 local_ordinal as integer {0..8})]];380 local_ordinal = (local_ordinal + 1) as integer {0..6};381 if left_scale_present then382 left_scale = _Tiles[[BundleMatrixSourceAt(383 local_ordinal as integer {0..8})]];384 local_ordinal = (local_ordinal + 1) as integer {0..6};385 end;386 right = _Tiles[[BundleMatrixSourceAt(387 local_ordinal as integer {0..8})]];388 local_ordinal = (local_ordinal + 1) as integer {0..6};389 if right_scale_present then390 right_scale = _Tiles[[BundleMatrixSourceAt(391 local_ordinal as integer {0..8})]];392 local_ordinal = (local_ordinal + 1) as integer {0..6};393 end;394 else395 let right_group = TileMatrixRightGroupSourceCount(396 function, right_type);397 if shared_count == right_group then398 left = _Tiles[[BundleMatrixSourceAt(399 local_ordinal as integer {0..8})]];400 local_ordinal = (local_ordinal + 1) as integer {0..6};401 if left_scale_present then402 left_scale = _Tiles[[BundleMatrixSourceAt(403 local_ordinal as integer {0..8})]];404 local_ordinal = (local_ordinal + 1) as integer {0..6};405 end;406 else407 left = MaterializeBundleSharedMatrixLeftPrimary(408 shared_ordinal as integer {0..3},409 m, k, left_type,410 _BundleFixedPointAttributes.trans_a,411 _CurrentMemoryAgent);412 shared_ordinal = (shared_ordinal + 1) as integer {0..4};413 if left_scale_present then414 left_scale = MaterializeBundleSharedMatrixLeftScale(415 shared_ordinal as integer {0..3},416 m, k, left_type,417 _BundleFixedPointAttributes.trans_a,418 _CurrentMemoryAgent);419 shared_ordinal = (shared_ordinal + 1) as integer {0..4};420 end;421 end;422 right = MaterializeBundleSharedMatrixPrimary(423 shared_ordinal as integer {0..3},424 k, n, right_type,425 _BundleFixedPointAttributes.trans_b,426 _CurrentMemoryAgent);427 shared_ordinal = (shared_ordinal + 1) as integer {0..4};428 if right_scale_present then429 let scale_groups = TileMXScaleGroupCount(k, right_type);430 right_scale = MaterializeBundleSharedMatrixPrimary(431 shared_ordinal as integer {0..3},432 scale_groups, n,433 TileMXScaleCarrierType(right_type),434 _BundleFixedPointAttributes.trans_b,435 _CurrentMemoryAgent);436 shared_ordinal = (shared_ordinal + 1) as integer {0..4};437 end;438 end;439
440 if TileMatrixFunctionUsesBias(function) then441 bias = BundleMatrixSourceAt(442 local_ordinal as integer {0..8});443 local_ordinal = (local_ordinal + 1) as integer {0..6};444 end;445
446 if _BundleFixedPointAttributes.c_scale_en then447 c_scale = BundleMatrixSourceAt(448 local_ordinal as integer {0..8});449 end;450
451 let right_group = TileMatrixRightGroupSourceCount(452 function, right_type);453 let shape_legal = if shared_count == 0 then454 TileMatrixCubeInfosMatchDimensions(left, right, pe_m, n, k)455 else if shared_count == right_group then456 TileMatrixMixedInfosMatchDimensions(left, right, pe_m, n, k)457 else458 TileMatrixInfosMatchDimensions(left, right, pe_m, n, k);459 let operand_types_legal = left.data_type == left_type &&460 right.data_type == right_type;461 let scales_legal = !TileMatrixFunctionUsesMX(function) ||462 TileMatrixInfoOptionalScalesLegal(463 left, left_scale, left_scale_present,464 right, right_scale, right_scale_present);465 assert shape_legal && operand_types_legal && scales_legal;466
467 let accumulator_legal = !TileMatrixFunctionUsesAccumulator(function) ||468 TileMatrixLocalCubeAccumulatorSchemaLegal(469 accumulator, pe_m, n, result_type, primary_layout,470 BundleMatrixPrimaryDestinationCapacityBytes());471 assert accumulator_legal;472 assert !TileMatrixFunctionUsesBias(function) ||473 TileMatrixInfoBiasLegal(474 left, right, bias, TileMatrixFunctionUsesMX(function));475 let destination = BundleMatrixDestinationAt(0);476 if TileMatrixFunctionUsesMX(function) then477 TMATMULMXSharedWithOptionalScales(478 destination, accumulator,479 left, left_scale, left_scale_present,480 right, right_scale, right_scale_present,481 bias, TileMatrixFunctionUsesBias(function),482 TileMatrixFunctionUsesAccumulator(function),483 c_scale, _BundleFixedPointAttributes.c_scale_en);484 else485 TMATMULShared(486 destination, accumulator, left, right, bias,487 TileMatrixFunctionUsesBias(function),488 TileMatrixFunctionUsesAccumulator(function),489 c_scale, _BundleFixedPointAttributes.c_scale_en);490 end;491 if _LastFault != Fault_None then492 RollBackBundleTileDestinations();493 return FALSE;494 end;495 if shared_count > 0 then496 ConsumeBundleSharedBindings(shared_count as integer {1..4});497 end;498 FinalizeBundleTileAttempt(TileExecution_Executed);499 return TRUE;500end;501