blazingly faster
Validate Operations / validate-operations (push) Has been cancelled

This commit is contained in:
NiccoloN
2026-07-19 09:59:49 +02:00
parent 5f42da36ae
commit ab54243fda
76 changed files with 4363 additions and 4323 deletions
+51 -10
View File
@@ -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());