Unexpected invariant now it's clear (batched in the first tensor rank)
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
ilgeco
2026-07-13 12:05:59 +02:00
parent fed6d343e5
commit 61e3ea9996
29 changed files with 2791 additions and 707 deletions
@@ -19,9 +19,7 @@ RankedTensorType getRowStripFragmentType(RankedTensorType logicalType) {
}
RankedTensorType getRowStripStorageType(RankedTensorType logicalType) {
return RankedTensorType::get({logicalType.getDimSize(2), logicalType.getDimSize(1), 1, logicalType.getDimSize(3)},
logicalType.getElementType(),
logicalType.getEncoding());
return spatial::getGraphBatchPhysicalResultType(logicalType.getDimSize(2), getRowStripFragmentType(logicalType));
}
std::pair<SmallVector<int64_t>, SmallVector<int64_t>> buildRowStripMetadata(RankedTensorType type) {
@@ -39,29 +37,12 @@ std::pair<SmallVector<int64_t>, SmallVector<int64_t>> buildRowStripMetadata(Rank
return {offsets, sizes};
}
SmallVector<OpFoldResult> buildRowStripFragmentOffsets(PatternRewriter& rewriter, OpFoldResult row) {
return {row, rewriter.getIndexAttr(0), rewriter.getIndexAttr(0), rewriter.getIndexAttr(0)};
}
SmallVector<OpFoldResult> buildRowStripFragmentSizes(PatternRewriter& rewriter, RankedTensorType logicalType) {
return {rewriter.getIndexAttr(1),
rewriter.getIndexAttr(logicalType.getDimSize(1)),
rewriter.getIndexAttr(1),
rewriter.getIndexAttr(logicalType.getDimSize(3))};
}
Value extractRowStripFragment(Value storage,
RankedTensorType logicalType,
OpFoldResult row,
PatternRewriter& rewriter,
Location loc) {
return tensor::ExtractSliceOp::create(rewriter,
loc,
getRowStripFragmentType(logicalType),
storage,
buildRowStripFragmentOffsets(rewriter, row),
buildRowStripFragmentSizes(rewriter, logicalType),
getUnitStrides(rewriter, 4));
return *extractGraphBatchPhysicalFragment(rewriter, loc, storage, row, getRowStripFragmentType(logicalType));
}
void insertRowStripFragment(Value fragment,
@@ -70,13 +51,11 @@ void insertRowStripFragment(Value fragment,
OpFoldResult row,
PatternRewriter& rewriter,
Location loc) {
createParallelInsertSliceIntoBatchOutput(rewriter,
loc,
fragment,
output,
buildRowStripFragmentOffsets(rewriter, row),
buildRowStripFragmentSizes(rewriter, logicalType),
getUnitStrides(rewriter, 4));
assert(fragment.getType() == getRowStripFragmentType(logicalType));
assert(output.getType() == getRowStripStorageType(logicalType));
auto slot = dyn_cast<Value>(row);
assert(slot && "row-strip graph publication requires a dynamic physical slot");
publishGraphBatchPhysicalFragment(rewriter, loc, fragment, output, slot);
}
FailureOr<Value> createPerChannelConstantFragment(DenseElementsAttr denseAttr,
@@ -145,30 +124,23 @@ FailureOr<Value> createRowStripStorageFromRows(Value rows,
}
FailureOr<Value>
materializeRowStripStorageToDense(Value storage, RankedTensorType logicalType, PatternRewriter& rewriter, Location loc) {
createRowStripAssemblyBlueprint(Value storage, RankedTensorType logicalType, PatternRewriter& rewriter, Location loc) {
auto storageType = dyn_cast<RankedTensorType>(storage.getType());
if (!storageType || storageType != getRowStripStorageType(logicalType))
return failure();
auto batchOp = createSpatComputeBatch(
rewriter, loc, TypeRange {logicalType}, logicalType.getDimSize(2), {}, ValueRange {storage},
[&](detail::SpatComputeBatchBodyArgs args) {
Value fragment = extractRowStripFragment(args.inputs.front(), logicalType, args.lane, rewriter, loc);
createParallelInsertSliceIntoBatchOutput(rewriter,
loc,
fragment,
args.outputs.front(),
SmallVector<OpFoldResult> {rewriter.getIndexAttr(0),
rewriter.getIndexAttr(0),
args.lane,
rewriter.getIndexAttr(0)},
buildRowStripFragmentSizes(rewriter, logicalType),
getUnitStrides(rewriter, 4));
return success();
});
if (failed(batchOp))
return failure();
return batchOp->getResult(0);
auto [offsets, sizes] = buildRowStripMetadata(logicalType);
int64_t height = logicalType.getDimSize(2);
SmallVector<int64_t> operandIndices(height, 0), sourceSlots, sourceOffsets(height, 0), strides(height * 4, 1);
for (int64_t row = 0; row < height; ++row)
sourceSlots.push_back(row);
return spatial::SpatBlueprintOp::create(rewriter, loc, logicalType, storage, ValueRange {},
rewriter.getStringAttr("nchw"), rewriter.getStringAttr("nchw_row_strip"),
rewriter.getDenseI64ArrayAttr(offsets), rewriter.getDenseI64ArrayAttr(sizes),
rewriter.getStringAttr("nchw_row_strip_fragments"), rewriter.getStringAttr("fragment_assembly"),
rewriter.getDenseI64ArrayAttr(operandIndices), rewriter.getDenseI64ArrayAttr(sourceSlots),
rewriter.getDenseI64ArrayAttr(sourceOffsets), rewriter.getDenseI64ArrayAttr(strides),
rewriter.getStringAttr("disjoint"), rewriter.getStringAttr("complete")).getOutput();
}
FailureOr<Value>