fix much stuff
This commit is contained in:
@@ -432,15 +432,6 @@ LogicalResult collectHostOutputs(MaterializerState& state) {
|
||||
return success();
|
||||
}
|
||||
|
||||
void setOperandSegmentSizes(Operation* op, int weightCount, int inputCount) {
|
||||
if (auto compute = dyn_cast<SpatCompute>(op)) {
|
||||
compute.getProperties().setOperandSegmentSizes({weightCount, inputCount});
|
||||
return;
|
||||
}
|
||||
auto batch = cast<SpatComputeBatch>(op);
|
||||
batch.getProperties().setOperandSegmentSizes({weightCount, inputCount});
|
||||
}
|
||||
|
||||
void createEmptyMaterializedOps(MaterializerState& state) {
|
||||
Location loc = state.func.getLoc();
|
||||
Block& funcBlock = state.func.getBody().front();
|
||||
@@ -529,19 +520,17 @@ BlockArgument appendWeight(MaterializerState& state, MaterializedClass& material
|
||||
materializedClass.weights.push_back(weight);
|
||||
|
||||
if (auto compute = dyn_cast<SpatCompute>(materializedClass.op)) {
|
||||
compute.getWeightsMutable().append(ValueRange(weight));
|
||||
setOperandSegmentSizes(materializedClass.op, materializedClass.weights.size(), materializedClass.inputs.size());
|
||||
BlockArgument arg = materializedClass.body->insertArgument(weightIndex, weight.getType(), weight.getLoc());
|
||||
materializedClass.weightArgs[weight] = arg;
|
||||
return arg;
|
||||
auto arg = compute.insertWeight(weightIndex, weight, weight.getLoc());
|
||||
assert(arg && "expected compute body while inserting a weight");
|
||||
materializedClass.weightArgs[weight] = std::get<1>(*arg);
|
||||
return std::get<1>(*arg);
|
||||
}
|
||||
|
||||
auto batch = cast<SpatComputeBatch>(materializedClass.op);
|
||||
batch.getWeightsMutable().append(ValueRange(weight));
|
||||
setOperandSegmentSizes(materializedClass.op, materializedClass.weights.size(), materializedClass.inputs.size());
|
||||
BlockArgument arg = materializedClass.body->insertArgument(1 + weightIndex, weight.getType(), weight.getLoc());
|
||||
materializedClass.weightArgs[weight] = arg;
|
||||
return arg;
|
||||
auto arg = batch.insertWeight(weightIndex, weight, weight.getLoc());
|
||||
assert(arg && "expected compute_batch body while inserting a weight argument");
|
||||
materializedClass.weightArgs[weight] = std::get<1>(*arg);
|
||||
return std::get<1>(*arg);
|
||||
}
|
||||
|
||||
BlockArgument appendInput(MaterializerState& state, MaterializedClass& materializedClass, Value input) {
|
||||
@@ -551,17 +540,16 @@ BlockArgument appendInput(MaterializerState& state, MaterializedClass& materiali
|
||||
|
||||
materializedClass.inputs.push_back(input);
|
||||
if (auto compute = dyn_cast<SpatCompute>(materializedClass.op)) {
|
||||
compute.getInputsMutable().append(ValueRange(input));
|
||||
BlockArgument arg = materializedClass.body->addArgument(input.getType(), input.getLoc());
|
||||
materializedClass.inputArgs[input] = arg;
|
||||
auto arg = compute.insertInput(materializedClass.inputs.size() - 1, input, input.getLoc());
|
||||
assert(arg && "expected compute body while inserting an input");
|
||||
materializedClass.inputArgs[input] = std::get<1>(*arg);
|
||||
return std::get<1>(*arg);
|
||||
}
|
||||
else {
|
||||
cast<SpatComputeBatch>(materializedClass.op).getInputsMutable().append(ValueRange(input));
|
||||
setOperandSegmentSizes(materializedClass.op, materializedClass.weights.size(), materializedClass.inputs.size());
|
||||
BlockArgument arg = materializedClass.body->insertArgument(
|
||||
materializedClass.body->getNumArguments() - 1, input.getType(), input.getLoc());
|
||||
materializedClass.inputArgs[input] = arg;
|
||||
return arg;
|
||||
if (auto compute = dyn_cast<SpatComputeBatch>(materializedClass.op)) {
|
||||
auto arg = compute.insertInput(materializedClass.inputs.size() - 1, input, input.getLoc());
|
||||
assert(arg && "expected compute_batch body while inserting an input argument");
|
||||
materializedClass.inputArgs[input] = std::get<1>(*arg);
|
||||
return std::get<1>(*arg);
|
||||
}
|
||||
llvm_unreachable("Cannot reach here");
|
||||
}
|
||||
@@ -608,6 +596,8 @@ Value createOriginalLaneValue(MaterializerState& state,
|
||||
return createIndexConstant(state, materializedClass.op, peers.front().laneStart);
|
||||
|
||||
auto batch = cast<SpatComputeBatch>(materializedClass.op);
|
||||
auto laneArg = batch.getLaneArgument();
|
||||
assert(laneArg && "expected materialized compute_batch lane argument");
|
||||
bool identity = true;
|
||||
for (auto [lane, peer] : llvm::enumerate(peers)) {
|
||||
if (peer.laneCount != 1 || peer.laneStart != lane) {
|
||||
@@ -616,7 +606,7 @@ Value createOriginalLaneValue(MaterializerState& state,
|
||||
}
|
||||
}
|
||||
if (identity)
|
||||
return batch.getLaneArgument();
|
||||
return *laneArg;
|
||||
|
||||
bool affineWithBase = true;
|
||||
int64_t base = static_cast<int64_t>(peers.front().laneStart);
|
||||
@@ -628,9 +618,9 @@ Value createOriginalLaneValue(MaterializerState& state,
|
||||
}
|
||||
if (affineWithBase) {
|
||||
if (base == 0)
|
||||
return batch.getLaneArgument();
|
||||
return *laneArg;
|
||||
Value baseValue = createIndexConstant(state, materializedClass.op, base);
|
||||
return arith::AddIOp::create(state.rewriter, loc, batch.getLaneArgument(), baseValue).getResult();
|
||||
return arith::AddIOp::create(state.rewriter, loc, *laneArg, baseValue).getResult();
|
||||
}
|
||||
|
||||
SmallVector<APInt, 8> laneValues;
|
||||
@@ -641,7 +631,7 @@ Value createOriginalLaneValue(MaterializerState& state,
|
||||
auto tableType = RankedTensorType::get({static_cast<int64_t>(peers.size())}, state.rewriter.getIndexType());
|
||||
auto tableAttr = DenseIntElementsAttr::get(tableType, laneValues);
|
||||
Value table = arith::ConstantOp::create(state.rewriter, loc, tableType, tableAttr).getResult();
|
||||
return tensor::ExtractOp::create(state.rewriter, loc, table, ValueRange {batch.getLaneArgument()}).getResult();
|
||||
return tensor::ExtractOp::create(state.rewriter, loc, table, ValueRange {*laneArg}).getResult();
|
||||
}
|
||||
|
||||
bool hasLiveExternalUse(Value value, const DenseSet<Operation*>& oldComputeOps) {
|
||||
@@ -838,7 +828,10 @@ setHostOutputValue(MaterializerState& state, MaterializedClass& sourceClass, Val
|
||||
offsets.reserve(payloadType.getRank());
|
||||
sizes.reserve(payloadType.getRank());
|
||||
strides.reserve(payloadType.getRank());
|
||||
offsets.push_back(batch.getLaneArgument());
|
||||
auto laneArg = batch.getLaneArgument();
|
||||
if (!laneArg)
|
||||
return batch.emitOpError("expected compute_batch lane block argument while materializing batch output");
|
||||
offsets.push_back(*laneArg);
|
||||
sizes.push_back(state.rewriter.getIndexAttr(1));
|
||||
strides.push_back(state.rewriter.getIndexAttr(1));
|
||||
for (int64_t dim = 1; dim < payloadType.getRank(); ++dim) {
|
||||
@@ -847,8 +840,11 @@ setHostOutputValue(MaterializerState& state, MaterializedClass& sourceClass, Val
|
||||
strides.push_back(state.rewriter.getIndexAttr(1));
|
||||
}
|
||||
|
||||
tensor::ParallelInsertSliceOp::create(
|
||||
state.rewriter, payload.getLoc(), payload, batch.getOutputArgument(resultIndex), offsets, sizes, strides);
|
||||
auto outputArg = batch.getOutputArgument(resultIndex);
|
||||
if (!outputArg)
|
||||
return batch.emitOpError("expected compute_batch output block argument while materializing batch output");
|
||||
|
||||
tensor::ParallelInsertSliceOp::create(state.rewriter, payload.getLoc(), payload, *outputArg, offsets, sizes, strides);
|
||||
return success();
|
||||
}
|
||||
|
||||
@@ -1136,14 +1132,20 @@ void mapWeights(MaterializerState& state,
|
||||
IRMapping& mapper) {
|
||||
Operation* op = instance.op;
|
||||
if (auto compute = dyn_cast<SpatCompute>(op)) {
|
||||
for (auto [index, weight] : llvm::enumerate(compute.getWeights()))
|
||||
mapper.map(compute.getWeightArgument(index), appendWeight(state, targetClass, weight));
|
||||
for (auto [index, weight] : llvm::enumerate(compute.getWeights())) {
|
||||
auto weightArg = compute.getWeightArgument(index);
|
||||
assert(weightArg && "expected compute weight block argument");
|
||||
mapper.map(*weightArg, appendWeight(state, targetClass, weight));
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
auto batch = cast<SpatComputeBatch>(op);
|
||||
for (auto [index, weight] : llvm::enumerate(batch.getWeights()))
|
||||
mapper.map(batch.getWeightArgument(index), appendWeight(state, targetClass, weight));
|
||||
for (auto [index, weight] : llvm::enumerate(batch.getWeights())) {
|
||||
auto weightArg = batch.getWeightArgument(index);
|
||||
assert(weightArg && "expected compute_batch weight block argument");
|
||||
mapper.map(*weightArg, appendWeight(state, targetClass, weight));
|
||||
}
|
||||
}
|
||||
|
||||
LogicalResult mapInputs(MaterializerState& state,
|
||||
@@ -1156,7 +1158,10 @@ LogicalResult mapInputs(MaterializerState& state,
|
||||
FailureOr<Value> mapped = resolveInputValue(state, targetClass, input, instance);
|
||||
if (failed(mapped))
|
||||
return compute.emitOpError("failed to resolve materialized compute input");
|
||||
mapper.map(compute.getInputArgument(index), *mapped);
|
||||
auto inputArg = compute.getInputArgument(index);
|
||||
if (!inputArg)
|
||||
return compute.emitOpError("expected compute input block argument while materializing inputs");
|
||||
mapper.map(*inputArg, *mapped);
|
||||
}
|
||||
return success();
|
||||
}
|
||||
@@ -1166,7 +1171,10 @@ LogicalResult mapInputs(MaterializerState& state,
|
||||
FailureOr<Value> mapped = resolveInputValue(state, targetClass, input, instance);
|
||||
if (failed(mapped))
|
||||
return batch.emitOpError("failed to resolve materialized compute_batch input");
|
||||
mapper.map(batch.getInputArgument(index), *mapped);
|
||||
auto inputArg = batch.getInputArgument(index);
|
||||
if (!inputArg)
|
||||
return batch.emitOpError("expected compute_batch input block argument while materializing inputs");
|
||||
mapper.map(*inputArg, *mapped);
|
||||
}
|
||||
return success();
|
||||
}
|
||||
@@ -1186,8 +1194,10 @@ SmallVector<Value, 4> collectMappedBatchOutputs(SpatComputeBatch batch, IRMappin
|
||||
if (!outputArg || outputArg.getOwner() != &batch.getBody().front())
|
||||
continue;
|
||||
|
||||
unsigned firstOutputArg = batch.getOutputArgument(0).getArgNumber();
|
||||
unsigned resultIndex = outputArg.getArgNumber() - firstOutputArg;
|
||||
auto firstOutputArg = batch.getOutputArgument(0);
|
||||
if (!firstOutputArg)
|
||||
return outputs;
|
||||
unsigned resultIndex = outputArg.getArgNumber() - firstOutputArg->getArgNumber();
|
||||
if (resultIndex >= outputs.size())
|
||||
continue;
|
||||
outputs[resultIndex] = mapper.lookupOrDefault(insert.getSource());
|
||||
@@ -1217,7 +1227,12 @@ cloneInstanceBody(MaterializerState& state, MaterializedClass& targetClass, Arra
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
mapper.map(batch.getLaneArgument(), createOriginalLaneValue(state, targetClass, peers, loc));
|
||||
auto laneArg = batch.getLaneArgument();
|
||||
if (!laneArg) {
|
||||
sourceOp->emitError("expected source compute_batch lane block argument");
|
||||
return failure();
|
||||
}
|
||||
mapper.map(*laneArg, createOriginalLaneValue(state, targetClass, peers, loc));
|
||||
}
|
||||
|
||||
mapWeights(state, targetClass, instance, mapper);
|
||||
|
||||
@@ -223,18 +223,32 @@ void mergeTriviallyConnectedComputes(func::FuncOp funcOp) {
|
||||
newBody->addArgument(input.getType(), loc);
|
||||
|
||||
IRMapping mapper;
|
||||
for (auto [weightIndex, _] : llvm::enumerate(compute.getWeights()))
|
||||
mapper.map(compute.getWeightArgument(weightIndex), newCompute.getWeightArgument(weightIndex));
|
||||
for (auto [inputIndex, _] : llvm::enumerate(compute.getInputs()))
|
||||
mapper.map(compute.getInputArgument(inputIndex), newCompute.getInputArgument(inputIndex));
|
||||
for (auto [oldIndex, weight] : llvm::enumerate(child.getWeights()))
|
||||
mapper.map(child.getWeightArgument(oldIndex), newCompute.getWeightArgument(childWeightToNewIndex[oldIndex]));
|
||||
for (auto [weightIndex, _] : llvm::enumerate(compute.getWeights())) {
|
||||
auto oldWeightArg = compute.getWeightArgument(weightIndex);
|
||||
auto newWeightArg = newCompute.getWeightArgument(weightIndex);
|
||||
assert(oldWeightArg && newWeightArg && "expected compute weight block arguments");
|
||||
mapper.map(*oldWeightArg, *newWeightArg);
|
||||
}
|
||||
for (auto [inputIndex, _] : llvm::enumerate(compute.getInputs())) {
|
||||
auto oldInputArg = compute.getInputArgument(inputIndex);
|
||||
auto newInputArg = newCompute.getInputArgument(inputIndex);
|
||||
assert(oldInputArg && newInputArg && "expected compute input block arguments");
|
||||
mapper.map(*oldInputArg, *newInputArg);
|
||||
}
|
||||
for (auto [oldIndex, weight] : llvm::enumerate(child.getWeights())) {
|
||||
auto oldWeightArg = child.getWeightArgument(oldIndex);
|
||||
auto newWeightArg = newCompute.getWeightArgument(childWeightToNewIndex[oldIndex]);
|
||||
assert(oldWeightArg && newWeightArg && "expected child compute weight block arguments");
|
||||
mapper.map(*oldWeightArg, *newWeightArg);
|
||||
}
|
||||
|
||||
rewriter.setInsertionPointToEnd(newBody);
|
||||
auto computeYield = cast<spatial::SpatYieldOp>(compute.getBody().front().getTerminator());
|
||||
for (Operation& op : compute.getBody().front().without_terminator())
|
||||
rewriter.clone(op, mapper);
|
||||
mapper.map(child.getInputArgument(childInputIndex), mapper.lookupOrDefault(computeYield.getOperand(usedResult)));
|
||||
auto childInputArg = child.getInputArgument(childInputIndex);
|
||||
assert(childInputArg && "expected child compute input block argument");
|
||||
mapper.map(*childInputArg, mapper.lookupOrDefault(computeYield.getOperand(usedResult)));
|
||||
|
||||
rewriter.setInsertionPointToEnd(newBody);
|
||||
for (auto& op : child.getBody().front())
|
||||
|
||||
Reference in New Issue
Block a user