This commit is contained in:
+923
-2846
File diff suppressed because it is too large
Load Diff
+9510
File diff suppressed because it is too large
Load Diff
+7548
File diff suppressed because it is too large
Load Diff
+128
@@ -0,0 +1,128 @@
|
||||
--- src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MaterializeMergeSchedule.cpp 2026-06-24 18:51:29.043731129 +0000
|
||||
+++ src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/MaterializeMergeSchedule.cpp 2026-06-24 18:51:29.026726895 +0000
|
||||
@@ -4112,104 +4112,8 @@
|
||||
Value originalOutput,
|
||||
Location loc);
|
||||
|
||||
-FailureOr<SmallVector<OpFoldResult, 4>> rematerializeProjectionIndexListForBatchHostOutput(
|
||||
- MaterializerState& state,
|
||||
- MaterializedClass& sourceClass,
|
||||
- ArrayRef<OpFoldResult> values,
|
||||
- IRMapping& mapper,
|
||||
- Location loc) {
|
||||
- SmallVector<OpFoldResult, 4> localized;
|
||||
- localized.reserve(values.size());
|
||||
- for (OpFoldResult value : values) {
|
||||
- FailureOr<OpFoldResult> remapped =
|
||||
- rematerializeIndexOpFoldResultInClass(state, sourceClass, value, loc, &mapper);
|
||||
- if (failed(remapped))
|
||||
- return failure();
|
||||
- localized.push_back(*remapped);
|
||||
- }
|
||||
- return localized;
|
||||
-}
|
||||
-
|
||||
-LogicalResult createProjectionAwareBatchHostInsert(MaterializerState& state,
|
||||
- MaterializedClass& sourceClass,
|
||||
- Value originalOutput,
|
||||
- Value payload,
|
||||
- Value destination,
|
||||
- ArrayRef<ProducerKey> keys,
|
||||
- Location loc) {
|
||||
- auto originalResult = dyn_cast<OpResult>(originalOutput);
|
||||
- if (!originalResult)
|
||||
- return failure();
|
||||
-
|
||||
- auto sourceBatch = dyn_cast_or_null<SpatComputeBatch>(originalResult.getOwner());
|
||||
- if (!sourceBatch || sourceBatch.getNumResults() == 0)
|
||||
- return failure();
|
||||
-
|
||||
- FailureOr<tensor::ParallelInsertSliceOp> projection =
|
||||
- getBatchResultProjectionInsert(sourceBatch, originalResult.getResultNumber());
|
||||
- if (failed(projection))
|
||||
- return failure();
|
||||
-
|
||||
- auto sourceLaneArg = sourceBatch.getLaneArgument();
|
||||
- if (!sourceLaneArg)
|
||||
- return failure();
|
||||
-
|
||||
- auto materializedBatch = dyn_cast<SpatScheduledComputeBatch>(sourceClass.op);
|
||||
- if (!materializedBatch)
|
||||
- return failure();
|
||||
-
|
||||
- auto materializedLaneArg = materializedBatch.getLaneArgument();
|
||||
- if (!materializedLaneArg)
|
||||
- return failure();
|
||||
-
|
||||
- if (keys.size() != sourceClass.cpus.size())
|
||||
- return failure();
|
||||
-
|
||||
- SmallVector<int64_t, 8> logicalLanes;
|
||||
- logicalLanes.reserve(keys.size());
|
||||
- for (ProducerKey key : keys) {
|
||||
- if (key.instance.op != sourceBatch.getOperation() || key.resultIndex != originalResult.getResultNumber())
|
||||
- return failure();
|
||||
- logicalLanes.push_back(key.instance.laneStart);
|
||||
- }
|
||||
-
|
||||
- IRMapping mapper;
|
||||
- Value logicalLane = createIndexedIndexValue(state,
|
||||
- sourceClass.op,
|
||||
- ArrayRef<int64_t>(logicalLanes),
|
||||
- *materializedLaneArg,
|
||||
- loc,
|
||||
- static_cast<int64_t>(sourceClass.cpus.size()),
|
||||
- /*allowExhaustiveTiledSearch=*/false);
|
||||
- mapper.map(*sourceLaneArg, logicalLane);
|
||||
-
|
||||
- FailureOr<SmallVector<OpFoldResult, 4>> offsets =
|
||||
- rematerializeProjectionIndexListForBatchHostOutput(
|
||||
- state, sourceClass, projection->getMixedOffsets(), mapper, loc);
|
||||
- if (failed(offsets))
|
||||
- return failure();
|
||||
- FailureOr<SmallVector<OpFoldResult, 4>> sizes =
|
||||
- rematerializeProjectionIndexListForBatchHostOutput(
|
||||
- state, sourceClass, projection->getMixedSizes(), mapper, loc);
|
||||
- if (failed(sizes))
|
||||
- return failure();
|
||||
- FailureOr<SmallVector<OpFoldResult, 4>> strides =
|
||||
- rematerializeProjectionIndexListForBatchHostOutput(
|
||||
- state, sourceClass, projection->getMixedStrides(), mapper, loc);
|
||||
- if (failed(strides))
|
||||
- return failure();
|
||||
-
|
||||
- tensor::ParallelInsertSliceOp::create(
|
||||
- state.rewriter, loc, payload, destination, *offsets, *sizes, *strides);
|
||||
- return success();
|
||||
-}
|
||||
-
|
||||
LogicalResult
|
||||
-setHostOutputValue(MaterializerState& state,
|
||||
- MaterializedClass& sourceClass,
|
||||
- Value originalOutput,
|
||||
- Value payload,
|
||||
- ArrayRef<ProducerKey> keys = {}) {
|
||||
+setHostOutputValue(MaterializerState& state, MaterializedClass& sourceClass, Value originalOutput, Value payload) {
|
||||
auto resultIt = sourceClass.hostOutputToResultIndex.find(originalOutput);
|
||||
if (resultIt == sourceClass.hostOutputToResultIndex.end())
|
||||
return sourceClass.op->emitError("missing host result slot for materialized output")
|
||||
@@ -4253,10 +4157,6 @@
|
||||
return batch.emitOpError("expected compute_batch output block argument while materializing batch output");
|
||||
|
||||
state.rewriter.setInsertionPointToStart(&inParallelOp.getRegion().front());
|
||||
- if (succeeded(createProjectionAwareBatchHostInsert(
|
||||
- state, sourceClass, originalOutput, payload, *outputArg, keys, payload.getLoc())))
|
||||
- return success();
|
||||
-
|
||||
createDim0ParallelInsertSlice(state, payload.getLoc(), payload, *outputArg, *laneArg);
|
||||
return success();
|
||||
}
|
||||
@@ -4276,7 +4176,7 @@
|
||||
|
||||
MaterializedClass& ownerClass = state.classes[ownerIt->second];
|
||||
if (sourceClass.id == ownerClass.id)
|
||||
- return setHostOutputValue(state, ownerClass, originalOutput, payload, keys);
|
||||
+ return setHostOutputValue(state, ownerClass, originalOutput, payload);
|
||||
|
||||
// Keep the old deadlock-free communication discipline: only scalar-to-scalar
|
||||
// host-owner forwarding is introduced here. Batch host publication remains on
|
||||
Reference in New Issue
Block a user