This commit is contained in:
@@ -457,39 +457,80 @@ void SpatDeferredCommunicationOp::print(OpAsmPrinter& printer) {
|
||||
printCompressedValueSequence(printer, getSources());
|
||||
printer.printOptionalAttrDict((*this)->getAttrs());
|
||||
printer << " : ";
|
||||
printer.printFunctionalType(getSources().getTypes(), getOperation()->getResultTypes());
|
||||
printCompressedTypeList(
|
||||
printer, getSources().getTypes(), ListDelimiter::Paren);
|
||||
printer << " -> ";
|
||||
printCompressedTypeSequence(printer, getOperation()->getResultTypes());
|
||||
printer << " ";
|
||||
printer.printRegion(getBody(), /*printEntryBlockArgs=*/false);
|
||||
}
|
||||
|
||||
ParseResult SpatDeferredCommunicationOp::parse(OpAsmParser& parser, OperationState& result) {
|
||||
SmallVector<OpAsmParser::UnresolvedOperand> sources;
|
||||
Type functionTypeStorage;
|
||||
SmallVector<Type> sourceTypes, outputTypes;
|
||||
|
||||
if (parseCompressedOperandSequence(parser, sources) || parser.parseOptionalAttrDict(result.attributes)
|
||||
|| parser.parseColon() || parser.parseType(functionTypeStorage))
|
||||
|| parser.parseColon()
|
||||
|| parseCompressedRepeatedList(
|
||||
parser, ListDelimiter::Paren, sourceTypes,
|
||||
[&](Type& type) { return parser.parseType(type); })
|
||||
|| parser.parseArrow()
|
||||
|| parseCompressedTypeSequence(
|
||||
parser, outputTypes, /*allowEmpty=*/false))
|
||||
return failure();
|
||||
|
||||
auto functionType = dyn_cast<FunctionType>(functionTypeStorage);
|
||||
if (!functionType)
|
||||
return parser.emitError(parser.getCurrentLocation(), "expected deferred communication function type");
|
||||
if (sources.size() != functionType.getNumInputs())
|
||||
if (sources.size() != sourceTypes.size())
|
||||
return parser.emitError(parser.getCurrentLocation(), "number of sources and source types must match");
|
||||
|
||||
if (parser.resolveOperands(sources, functionType.getInputs(), parser.getCurrentLocation(), result.operands))
|
||||
if (parser.resolveOperands(sources, sourceTypes, parser.getCurrentLocation(), result.operands))
|
||||
return failure();
|
||||
result.addTypes(functionType.getResults());
|
||||
result.addTypes(outputTypes);
|
||||
|
||||
Region* body = result.addRegion();
|
||||
SmallVector<OpAsmParser::Argument> bodyArgs;
|
||||
for (Type type : functionType.getInputs()) {
|
||||
for (Type type : sourceTypes) {
|
||||
OpAsmParser::Argument argument;
|
||||
argument.type = type;
|
||||
bodyArgs.push_back(argument);
|
||||
}
|
||||
if (auto count = dyn_cast_or_null<IntegerAttr>(
|
||||
result.attributes.get("specialization_count"));
|
||||
count && count.getInt() > 1) {
|
||||
OpAsmParser::Argument argument;
|
||||
argument.type = parser.getBuilder().getIndexType();
|
||||
bodyArgs.push_back(argument);
|
||||
}
|
||||
return parser.parseRegion(*body, bodyArgs);
|
||||
}
|
||||
|
||||
void SpatDeferredSourceSelectOp::print(OpAsmPrinter& printer) {
|
||||
printer << " " << getSelector() << " of ";
|
||||
printCompressedValueSequence(printer, getSources());
|
||||
printer.printOptionalAttrDict((*this)->getAttrs());
|
||||
printer << " : " << getOutput().getType();
|
||||
}
|
||||
|
||||
ParseResult SpatDeferredSourceSelectOp::parse(
|
||||
OpAsmParser& parser, OperationState& result) {
|
||||
OpAsmParser::UnresolvedOperand selector;
|
||||
SmallVector<OpAsmParser::UnresolvedOperand> sources;
|
||||
Type outputType;
|
||||
if (parser.parseOperand(selector) || parser.parseKeyword("of")
|
||||
|| parseCompressedOperandSequence(parser, sources)
|
||||
|| parser.parseOptionalAttrDict(result.attributes)
|
||||
|| parser.parseColon() || parser.parseType(outputType))
|
||||
return failure();
|
||||
if (parser.resolveOperand(selector, parser.getBuilder().getIndexType(),
|
||||
result.operands))
|
||||
return failure();
|
||||
SmallVector<Type> sourceTypes(sources.size(), outputType);
|
||||
if (parser.resolveOperands(sources, sourceTypes,
|
||||
parser.getCurrentLocation(), result.operands))
|
||||
return failure();
|
||||
result.addTypes(outputType);
|
||||
return success();
|
||||
}
|
||||
|
||||
void SpatExtractRowsOp::print(OpAsmPrinter& printer) {
|
||||
printer << " ";
|
||||
printer.printOperand(getInput());
|
||||
|
||||
Reference in New Issue
Block a user