================
@@ -110,6 +111,112 @@ Operation
*cir::CIRDialect::materializeConstant(mlir::OpBuilder &builder,
mlir::cast<mlir::TypedAttr>(value));
}
+//===----------------------------------------------------------------------===//
+// Offload container helpers
+//===----------------------------------------------------------------------===//
+
+bool cir::isOffloadContainer(mlir::ModuleOp module) {
+ return module->hasAttr(cir::CIRDialect::getOffloadContainerAttrName());
+}
+
+mlir::ModuleOp cir::getOffloadHostModule(mlir::ModuleOp container) {
+ assert(isOffloadContainer(container) && "expected an offload container");
+ return mlir::cast<mlir::ModuleOp>(container.getBody()->front());
+}
+
+llvm::iterator_range<mlir::Block::op_iterator<mlir::ModuleOp>>
+cir::getOffloadDeviceModules(mlir::ModuleOp container) {
+ assert(isOffloadContainer(container) && "expected an offload container");
+ mlir::Block &body = *container.getBody();
+ auto begin = body.op_begin<mlir::ModuleOp>();
+ auto end = body.op_end<mlir::ModuleOp>();
+ if (begin != end)
+ ++begin;
+ return {begin, end};
+}
+
+//===----------------------------------------------------------------------===//
+// Dialect attribute verification
+//===----------------------------------------------------------------------===//
+
+static LogicalResult verifyOffloadKind(mlir::ModuleOp module,
+ cir::OffloadKind expected) {
+ auto attr = module->getAttrOfType<cir::OffloadKindAttr>(
+ cir::CIRDialect::getOffloadKindAttrName());
+ if (!attr)
+ return module.emitOpError()
+ << "expects '" << cir::CIRDialect::getOffloadKindAttrName()
+ << "' offload kind attribute";
+ if (attr.getValue() != expected)
+ return module.emitOpError()
+ << "expects '" << cir::CIRDialect::getOffloadKindAttrName()
+ << "' value '" << cir::stringifyOffloadKind(expected) << "'";
+ return success();
+}
+
+// A module marked with `cir.offload.container` holds the host module followed
+// by one or more device modules, each tagged with `cir.offload.kind`. Keeping
+// the host module first gives later offload passes a simple convention for
+// finding the host side while iterating the remaining device modules.
+static LogicalResult verifyOffloadContainer(mlir::Operation *op) {
+ auto container = mlir::dyn_cast<mlir::ModuleOp>(op);
+ if (!container)
+ return op->emitError() << "expects '"
+ << cir::CIRDialect::getOffloadContainerAttrName()
+ << "' attribute to be attached to '"
+ << mlir::ModuleOp::getOperationName() << "'";
+
+ mlir::Block &body = *container.getBody();
+ if (body.empty())
+ return container.emitOpError()
+ << "expects host module as the first nested op";
+
+ auto host = mlir::dyn_cast<mlir::ModuleOp>(body.front());
+ if (!host)
+ return container.emitOpError()
+ << "expects host module as the first nested op";
+ if (failed(verifyOffloadKind(host, cir::OffloadKind::Host)))
+ return failure();
+
+ unsigned numDevices = 0;
+ auto it = body.begin();
+ for (++it; it != body.end(); ++it) {
+ auto module = mlir::dyn_cast<mlir::ModuleOp>(*it);
+ if (!module)
+ return container.emitOpError()
+ << "expects only nested builtin.module ops";
+ if (failed(verifyOffloadKind(module, cir::OffloadKind::Device)))
+ return failure();
+ ++numDevices;
+ }
+
+ if (numDevices == 0)
+ return container.emitOpError() << "expects at least one device module";
----------------
steffenlarsen wrote:
No need to count the device modules when we can determine it from the iterators.
```suggestion
auto it = body.begin() + 1;
if (it == body.end())
return container.emitOpError() << "expects at least one device module";
for (; it != body.end(); ++it) {
auto module = mlir::dyn_cast<mlir::ModuleOp>(*it);
if (!module)
return container.emitOpError()
<< "expects only nested builtin.module ops";
if (failed(verifyOffloadKind(module, cir::OffloadKind::Device)))
return failure();
}
```
https://github.com/llvm/llvm-project/pull/206576
_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits