19#include "mlir/IR/AttrTypeSubElements.h"
20#include "mlir/IR/Builders.h"
21#include "mlir/IR/BuiltinTypes.h"
22#include "mlir/IR/Diagnostics.h"
23#include "mlir/IR/DialectImplementation.h"
24#include "mlir/IR/StorageUniquerSupport.h"
25#include "mlir/IR/Types.h"
26#include "mlir/Interfaces/MemorySlotInterfaces.h"
27#include "llvm/ADT/SmallSet.h"
28#include "llvm/ADT/StringExtras.h"
29#include "llvm/ADT/StringSet.h"
30#include "llvm/ADT/TypeSwitch.h"
36static ParseResult
parseHWArray(AsmParser &parser, Attribute &dim,
43#define GET_TYPEDEF_CLASSES
44#include "circt/Dialect/HW/HWTypes.cpp.inc"
51 mlir::AttrTypeReplacer replacer;
52 replacer.addReplacement(
53 [](TypeAliasType alias) {
return alias.getCanonicalType(); });
54 return replacer.replace(type);
62 if (isa<hw::IntType>(canonicalType))
65 auto intType = dyn_cast<IntegerType>(canonicalType);
66 if (!intType || !intType.isSignless())
81 if (isa<IntegerType, IntType, EnumType>(type))
84 if (
auto array = dyn_cast<ArrayType>(type))
87 if (
auto array = dyn_cast<UnpackedArrayType>(type))
90 if (
auto t = dyn_cast<StructType>(type))
91 return llvm::all_of(t.getElements(),
92 [](
auto f) { return isHWValueType(f.type); });
94 if (
auto t = dyn_cast<UnionType>(type))
95 return llvm::all_of(t.getElements(),
96 [](
auto f) { return isHWValueType(f.type); });
98 if (
auto t = dyn_cast<TypeAliasType>(type))
108 if (isa<IntegerType>(type))
124 return llvm::TypeSwitch<::mlir::Type, int64_t>(type)
126 [](IntegerType t) {
return t.getIntOrFloatBitWidth(); })
127 .Default([](Type type) -> int64_t {
129 if (
auto iface = dyn_cast<BitWidthTypeInterface>(type)) {
130 std::optional<int64_t> width = iface.getBitWidth();
131 return width.has_value() ? *width : -1;
141 if (
auto array = dyn_cast<ArrayType>(type))
144 if (
auto array = dyn_cast<UnpackedArrayType>(type))
147 if (
auto t = dyn_cast<StructType>(type)) {
148 return std::any_of(t.getElements().begin(), t.getElements().end(),
149 [](
const auto &f) { return hasHWInOutType(f.type); });
152 if (
auto t = dyn_cast<TypeAliasType>(type))
155 return isa<InOutType>(type);
159struct AggregateAttrFrame {
160 SmallVector<Attribute> attrs;
161 SmallVector<Type> types;
164 AggregateAttrFrame(SmallVector<Type> &&types)
165 : attrs(types.size()), types(std::move(types)), remaining(attrs.size()) {}
167 void addChild(Attribute attr) { attrs[--remaining] = attr; }
168 Type getNextChildType() {
return types[remaining - 1]; }
169 bool isFinished()
const {
return remaining == 0; }
179 auto *ctx = aggregateType.getContext();
180 SmallVector<AggregateAttrFrame> stack;
181 unsigned nextExtraction = 0;
183 auto pushToStack = [&](Type type) ->
bool {
184 return TypeSwitch<Type, bool>(type)
185 .Case<StructType>([&](
auto structType) {
186 auto len = structType.getElements().size();
187 SmallVector<Type> types;
189 for (
auto &element : structType.getElements())
191 stack.push_back(std::move(types));
194 .Case<ArrayType, UnpackedArrayType>([&](
auto arrayType) {
195 SmallVector<Type> types(arrayType.getNumElements(),
197 stack.push_back(std::move(types));
209 while (!stack.empty()) {
210 if (stack.back().isFinished()) {
211 auto frame = stack.pop_back_val();
212 result = ArrayAttr::get(ctx, frame.attrs);
214 stack.back().addChild(result);
218 auto curType = stack.back().getNextChildType();
219 if (
auto intType = dyn_cast<IntegerType>(curType)) {
220 auto width = intType.getWidth();
221 auto elemValue = width ? intVal.extractBits(width, nextExtraction)
222 : APInt(0, 0,
false);
223 nextExtraction += width;
224 stack.back().addChild(IntegerAttr::get(intType, elemValue));
226 if (!pushToStack(curType))
231 assert(nextExtraction == intVal.getBitWidth() &&
232 "constant wasn't fully processed");
242 SmallVector<Attribute> worklist;
243 worklist.push_back(attr);
244 auto bitWidth = hw::getBitWidth(type);
245 assert(bitWidth >= 0 &&
"bit width must be known for constant");
246 result = APInt(bitWidth, 0);
247 unsigned nextInsertion = 0;
249 while (!worklist.empty()) {
250 auto current = worklist.pop_back_val();
251 if (
auto innerArray = dyn_cast<ArrayAttr>(current)) {
252 worklist.append(innerArray.begin(), innerArray.end());
256 if (
auto intAttr = dyn_cast<IntegerAttr>(current)) {
257 auto chunk = intAttr.getValue();
258 result.insertBits(chunk, nextInsertion);
259 nextInsertion += chunk.getBitWidth();
266 assert(nextInsertion == bitWidth &&
"constant wasn't fully processed");
276 auto fullString =
static_cast<DialectAsmParser &
>(p).getFullSymbolSpec();
277 auto *curPtr = p.getCurrentLocation().getPointer();
279 StringRef(curPtr, fullString.size() - (curPtr - fullString.data()));
281 if (typeString.starts_with(
"array<") || typeString.starts_with(
"inout<") ||
282 typeString.starts_with(
"uarray<") || typeString.starts_with(
"struct<") ||
283 typeString.starts_with(
"typealias<") || typeString.starts_with(
"int<") ||
284 typeString.starts_with(
"enum<") || typeString.starts_with(
"union<")) {
285 llvm::StringRef mnemonic;
286 if (
auto parseResult = generatedTypeParser(p, &mnemonic, result);
287 parseResult.has_value())
289 return p.emitError(p.getNameLoc(),
"invalid type `") << typeString <<
"`";
292 return p.parseType(result);
296 if (succeeded(generatedTypePrinter(element, p)))
298 p.printType(element);
305Type IntType::get(mlir::TypedAttr width) {
307 auto widthWidth = llvm::dyn_cast<IntegerType>(width.getType());
308 assert(widthWidth && widthWidth.getWidth() == 32 &&
309 "!hw.int width must be 32-bits");
312 if (
auto cstWidth = llvm::dyn_cast<IntegerAttr>(width))
313 return IntegerType::get(width.getContext(),
314 cstWidth.getValue().getZExtValue());
316 return Base::get(width.getContext(), width);
319Type IntType::parse(AsmParser &p) {
321 auto int32Type = p.getBuilder().getIntegerType(32);
323 mlir::TypedAttr width;
324 if (p.parseLess() || p.parseAttribute(width, int32Type) || p.parseGreater())
329void IntType::print(AsmPrinter &p)
const {
331 p.printAttributeWithoutType(
getWidth());
346 return llvm::hash_combine(fi.
name, fi.
type);
355 SmallVectorImpl<FieldInfo> ¶meters) {
356 llvm::StringSet<> nameSet;
357 bool hasDuplicateName =
false;
358 auto parseResult = p.parseCommaSeparatedList(
359 mlir::AsmParser::Delimiter::LessGreater, [&]() -> ParseResult {
363 auto fieldLoc = p.getCurrentLocation();
364 if (p.parseKeywordOrString(&name) || p.parseColon() ||
368 if (!nameSet.insert(name).second) {
369 p.emitError(fieldLoc,
"duplicate field name \'" + name +
"\'");
372 hasDuplicateName = true;
375 parameters.push_back(
376 FieldInfo{StringAttr::get(p.getContext(), name), type});
380 if (hasDuplicateName)
386static void printFields(AsmPrinter &p, ArrayRef<FieldInfo> fields) {
388 llvm::interleaveComma(fields, p, [&](
const FieldInfo &field) {
389 p.printKeywordOrString(field.
name.getValue());
390 p <<
": " << field.
type;
395Type StructType::parse(AsmParser &p) {
396 llvm::SmallVector<FieldInfo, 4> parameters;
399 return get(p.getContext(), parameters);
402LogicalResult StructType::verify(function_ref<InFlightDiagnostic()> emitError,
403 ArrayRef<StructType::FieldInfo> elements) {
404 llvm::SmallDenseSet<StringAttr> fieldNameSet;
405 LogicalResult result = success();
406 fieldNameSet.reserve(elements.size());
407 for (
const auto &elt : elements)
408 if (!fieldNameSet.insert(elt.name).second) {
410 emitError() <<
"duplicate field name '" << elt.name.getValue()
411 <<
"' in hw.struct type";
416void StructType::print(AsmPrinter &p)
const {
printFields(p, getElements()); }
418Type StructType::getFieldType(mlir::StringRef fieldName) {
419 for (
const auto &field : getElements())
420 if (field.name == fieldName)
425std::optional<uint32_t> StructType::getFieldIndex(mlir::StringRef fieldName) {
426 ArrayRef<hw::StructType::FieldInfo> elems = getElements();
427 for (
size_t idx = 0, numElems = elems.size(); idx < numElems; ++idx)
428 if (elems[idx].name == fieldName)
433std::optional<uint32_t> StructType::getFieldIndex(mlir::StringAttr fieldName) {
434 ArrayRef<hw::StructType::FieldInfo> elems = getElements();
435 for (
size_t idx = 0, numElems = elems.size(); idx < numElems; ++idx)
436 if (elems[idx].name == fieldName)
441static std::pair<uint64_t, SmallVector<uint64_t>>
443 uint64_t fieldID = 0;
444 auto elements = st.getElements();
445 SmallVector<uint64_t> fieldIDs;
446 fieldIDs.reserve(elements.size());
447 for (
auto &element : elements) {
448 auto type = element.type;
450 fieldIDs.push_back(fieldID);
454 return {fieldID, fieldIDs};
457void StructType::getInnerTypes(SmallVectorImpl<Type> &types) {
458 for (
const auto &field : getElements())
459 types.push_back(field.type);
462uint64_t StructType::getMaxFieldID()
const {
463 uint64_t fieldID = 0;
464 for (
const auto &field : getElements())
469std::pair<Type, uint64_t>
470StructType::getSubTypeByFieldID(uint64_t fieldID)
const {
474 auto *it = std::prev(llvm::upper_bound(fieldIDs, fieldID));
475 auto subfieldIndex = std::distance(fieldIDs.begin(), it);
476 auto subfieldType = getElements()[subfieldIndex].type;
477 auto subfieldID = fieldID - fieldIDs[subfieldIndex];
478 return {subfieldType, subfieldID};
481std::pair<uint64_t, bool>
482StructType::projectToChildFieldID(uint64_t fieldID, uint64_t index)
const {
484 auto childRoot = fieldIDs[index];
486 index + 1 >= getElements().size() ? maxId : (fieldIDs[index + 1] - 1);
487 return std::make_pair(fieldID - childRoot,
488 fieldID >= childRoot && fieldID <= rangeEnd);
491uint64_t StructType::getFieldID(uint64_t index)
const {
493 return fieldIDs[index];
496uint64_t StructType::getIndexForFieldID(uint64_t fieldID)
const {
497 assert(!getElements().
empty() &&
"Bundle must have >0 fields");
499 auto *it = std::prev(llvm::upper_bound(fieldIDs, fieldID));
500 return std::distance(fieldIDs.begin(), it);
503std::pair<uint64_t, uint64_t>
504StructType::getIndexAndSubfieldID(uint64_t fieldID)
const {
507 return {index, fieldID - elementFieldID};
510std::optional<DenseMap<Attribute, Type>>
511hw::StructType::getSubelementIndexMap()
const {
512 DenseMap<Attribute, Type> destructured;
513 for (
auto [i, field] :
llvm::enumerate(getElements()))
515 {IntegerAttr::get(IndexType::get(getContext()), i), field.type});
519Type hw::StructType::getTypeAtIndex(Attribute index)
const {
520 auto indexAttr = llvm::dyn_cast<IntegerAttr>(index);
527std::optional<int64_t> StructType::getBitWidth()
const {
529 for (
auto field : getElements()) {
530 int64_t fieldSize = hw::getBitWidth(field.type);
556Type UnionType::parse(AsmParser &p) {
557 llvm::SmallVector<FieldInfo, 4> parameters;
558 llvm::StringSet<> nameSet;
559 bool hasDuplicateName =
false;
560 if (p.parseCommaSeparatedList(
561 mlir::AsmParser::Delimiter::LessGreater, [&]() -> ParseResult {
565 auto fieldLoc = p.getCurrentLocation();
566 if (p.parseKeyword(&name) || p.parseColon() || p.parseType(type))
569 if (!nameSet.insert(name).second) {
570 p.emitError(fieldLoc,
"duplicate field name \'" + name +
571 "\' in hw.union type");
574 hasDuplicateName = true;
578 if (succeeded(p.parseOptionalKeyword(
"offset")))
579 if (p.parseInteger(offset))
581 parameters.push_back(UnionType::FieldInfo{
582 StringAttr::get(p.getContext(), name), type, offset});
587 if (hasDuplicateName)
590 return get(p.getContext(), parameters);
593void UnionType::print(AsmPrinter &odsPrinter)
const {
595 llvm::interleaveComma(
596 getElements(), odsPrinter, [&](
const UnionType::FieldInfo &field) {
597 odsPrinter << field.name.getValue() <<
": " << field.type;
599 odsPrinter <<
" offset " << field.offset;
604LogicalResult UnionType::verify(function_ref<InFlightDiagnostic()> emitError,
605 ArrayRef<UnionType::FieldInfo> elements) {
606 llvm::SmallDenseSet<StringAttr> fieldNameSet;
607 LogicalResult result = success();
608 fieldNameSet.reserve(elements.size());
609 for (
const auto &elt : elements)
610 if (!fieldNameSet.insert(elt.name).second) {
612 emitError() <<
"duplicate field name '" << elt.name.getValue()
613 <<
"' in hw.union type";
618std::optional<uint32_t> UnionType::getFieldIndex(mlir::StringAttr fieldName) {
619 ArrayRef<hw::UnionType::FieldInfo> elems = getElements();
620 for (
size_t idx = 0, numElems = elems.size(); idx < numElems; ++idx)
621 if (elems[idx].name == fieldName)
626std::optional<uint32_t> UnionType::getFieldIndex(mlir::StringRef fieldName) {
627 return getFieldIndex(StringAttr::get(getContext(), fieldName));
630UnionType::FieldInfo UnionType::getFieldInfo(::mlir::StringRef fieldName) {
631 if (
auto fieldIndex = getFieldIndex(fieldName))
632 return getElements()[*fieldIndex];
636Type UnionType::getFieldType(mlir::StringRef fieldName) {
637 return getFieldInfo(fieldName).type;
640std::optional<int64_t> UnionType::getBitWidth()
const {
642 for (
auto field : getElements()) {
643 int64_t fieldSize = hw::getBitWidth(field.type);
646 fieldSize += field.offset;
647 if (fieldSize > maxSize)
657Type EnumType::parse(AsmParser &p) {
658 llvm::SmallVector<Attribute> fields;
660 if (p.parseCommaSeparatedList(AsmParser::Delimiter::LessGreater, [&]() {
662 if (p.parseKeyword(&name))
664 fields.push_back(StringAttr::get(p.getContext(), name));
669 return get(p.getContext(), ArrayAttr::get(p.getContext(), fields));
672void EnumType::print(AsmPrinter &p)
const {
674 llvm::interleaveComma(getFields(), p, [&](Attribute enumerator) {
675 p << llvm::cast<StringAttr>(enumerator).getValue();
680bool EnumType::contains(mlir::StringRef field) {
681 return indexOf(field).has_value();
684std::optional<size_t> EnumType::indexOf(mlir::StringRef field) {
685 for (
auto it :
llvm::enumerate(getFields()))
686 if (
llvm::cast<StringAttr>(it.value()).getValue() == field)
691std::optional<int64_t> EnumType::getBitWidth()
const {
692 auto w = getFields().size();
694 return llvm::Log2_64_Ceil(w);
702static ParseResult
parseHWArray(AsmParser &p, Attribute &dim, Type &inner) {
704 auto int64Type = p.getBuilder().getIntegerType(64);
706 if (
auto res = p.parseOptionalInteger(dimLiteral); res.has_value()) {
709 dim = p.getBuilder().getI64IntegerAttr(dimLiteral);
710 }
else if (
auto res64 = p.parseOptionalAttribute(dim, int64Type);
715 return p.emitError(p.getNameLoc(),
"expected integer");
717 if (!isa<IntegerAttr, ParamExprAttr, ParamDeclRefAttr>(dim)) {
718 p.emitError(p.getNameLoc(),
"unsupported dimension kind in hw.array");
729 p.printAttributeWithoutType(dim);
734size_t ArrayType::getNumElements()
const {
735 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(getSizeAttr()))
736 return intAttr.getInt();
740LogicalResult ArrayType::verify(function_ref<InFlightDiagnostic()> emitError,
741 Type innerType, Attribute size) {
743 return emitError() <<
"hw.array cannot contain InOut types";
747uint64_t ArrayType::getMaxFieldID()
const {
748 return getNumElements() *
752std::pair<Type, uint64_t>
753ArrayType::getSubTypeByFieldID(uint64_t fieldID)
const {
759std::pair<uint64_t, bool>
760ArrayType::projectToChildFieldID(uint64_t fieldID, uint64_t index)
const {
764 return std::make_pair(fieldID - childRoot,
765 fieldID >= childRoot && fieldID <= rangeEnd);
768uint64_t ArrayType::getIndexForFieldID(uint64_t fieldID)
const {
769 assert(fieldID &&
"fieldID must be at least 1");
774std::pair<uint64_t, uint64_t>
775ArrayType::getIndexAndSubfieldID(uint64_t fieldID)
const {
778 return {index, fieldID - elementFieldID};
781uint64_t ArrayType::getFieldID(uint64_t index)
const {
785std::optional<DenseMap<Attribute, Type>>
786hw::ArrayType::getSubelementIndexMap()
const {
787 DenseMap<Attribute, Type> destructured;
788 for (
unsigned i = 0; i < getNumElements(); ++i)
790 {IntegerAttr::get(IndexType::get(getContext()), i), getElementType()});
794Type hw::ArrayType::getTypeAtIndex(Attribute index)
const {
795 return getElementType();
798std::optional<int64_t> hw::ArrayType::getBitWidth()
const {
799 auto elementBitWidth = hw::getBitWidth(getElementType());
800 if (elementBitWidth < 0)
813UnpackedArrayType::verify(function_ref<InFlightDiagnostic()> emitError,
814 Type innerType, Attribute size) {
816 return emitError() <<
"invalid element for uarray type";
820size_t UnpackedArrayType::getNumElements()
const {
821 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(getSizeAttr()))
822 return intAttr.getInt();
826uint64_t UnpackedArrayType::getMaxFieldID()
const {
827 return getNumElements() *
831std::pair<Type, uint64_t>
832UnpackedArrayType::getSubTypeByFieldID(uint64_t fieldID)
const {
838std::pair<uint64_t, bool>
839UnpackedArrayType::projectToChildFieldID(uint64_t fieldID,
840 uint64_t index)
const {
844 return std::make_pair(fieldID - childRoot,
845 fieldID >= childRoot && fieldID <= rangeEnd);
848uint64_t UnpackedArrayType::getIndexForFieldID(uint64_t fieldID)
const {
849 assert(fieldID &&
"fieldID must be at least 1");
854std::pair<uint64_t, uint64_t>
855UnpackedArrayType::getIndexAndSubfieldID(uint64_t fieldID)
const {
858 return {index, fieldID - elementFieldID};
861uint64_t UnpackedArrayType::getFieldID(uint64_t index)
const {
865std::optional<int64_t> UnpackedArrayType::getBitWidth()
const {
866 auto elementBitWidth = hw::getBitWidth(getElementType());
867 if (elementBitWidth < 0)
869 int64_t dimBitWidth = getNumElements();
872 return (int64_t)getNumElements() * elementBitWidth;
879LogicalResult InOutType::verify(function_ref<InFlightDiagnostic()> emitError,
882 return emitError() <<
"invalid element for hw.inout type " <<
innerType;
890TypeAliasType TypeAliasType::get(SymbolRefAttr ref, Type innerType) {
891 return get(ref.getContext(), ref, innerType, hw::getCanonicalType(innerType));
895TypeAliasType::getChecked(function_ref<InFlightDiagnostic()> emitError,
896 SymbolRefAttr ref, Type innerType) {
897 return getChecked(emitError, ref.getContext(), ref, innerType,
898 hw::getCanonicalType(innerType));
902TypeAliasType::verify(function_ref<InFlightDiagnostic()> emitError,
903 SymbolRefAttr ref, Type innerType, Type canonicalType) {
904 if (ref.getNestedReferences().size() != 1)
906 <<
"expected exactly one nested reference in hw.typealias";
910Type TypeAliasType::parse(AsmParser &p) {
913 if (p.parseLess() || p.parseAttribute(ref) || p.parseComma() ||
914 p.parseType(type) || p.parseGreater())
917 return p.getChecked<TypeAliasType>(ref, type);
920void TypeAliasType::print(AsmPrinter &p)
const {
921 p <<
"<" << getRef() <<
", " << getInnerType() <<
">";
926TypedeclOp TypeAliasType::getTypeDecl(
const HWSymbolCache &cache) {
927 SymbolRefAttr ref = getRef();
928 auto typeScope = ::dyn_cast_or_null<TypeScopeLike>(
933 return dyn_cast_or_null<TypedeclOp>(
934 SymbolTable::lookupSymbolIn(typeScope, ref.getLeafReference()));
937std::optional<int64_t> TypeAliasType::getBitWidth()
const {
948LogicalResult ModuleType::verify(function_ref<InFlightDiagnostic()> emitError,
949 ArrayRef<ModulePort> ports) {
950 if (llvm::any_of(ports, [](
const ModulePort &port) {
953 return emitError() <<
"Ports cannot be inout types";
957size_t ModuleType::getPortIdForInputId(
size_t idx) {
958 assert(idx < getImpl()->inputToAbs.size() &&
"input port out of range");
959 return getImpl()->inputToAbs[idx];
962size_t ModuleType::getPortIdForOutputId(
size_t idx) {
963 assert(idx < getImpl()->outputToAbs.size() &&
" output port out of range");
964 return getImpl()->outputToAbs[idx];
967size_t ModuleType::getInputIdForPortId(
size_t idx) {
968 auto nIdx = getImpl()->absToInput[idx];
973size_t ModuleType::getOutputIdForPortId(
size_t idx) {
974 auto nIdx = getImpl()->absToOutput[idx];
979size_t ModuleType::getNumInputs() {
return getImpl()->inputToAbs.size(); }
981size_t ModuleType::getNumOutputs() {
return getImpl()->outputToAbs.size(); }
983size_t ModuleType::getNumPorts() {
return getPorts().size(); }
985SmallVector<Type> ModuleType::getInputTypes() {
986 SmallVector<Type> retval;
987 for (
auto &p : getPorts()) {
988 if (p.dir == ModulePort::Direction::Input)
989 retval.push_back(p.type);
990 else if (p.dir == ModulePort::Direction::InOut) {
991 retval.push_back(hw::InOutType::get(p.type));
997SmallVector<Type> ModuleType::getOutputTypes() {
998 SmallVector<Type> retval;
999 for (
auto &p : getPorts())
1001 retval.push_back(p.type);
1005SmallVector<Type> ModuleType::getPortTypes() {
1006 SmallVector<Type> retval;
1007 for (
auto &p : getPorts())
1008 retval.push_back(p.type);
1012Type ModuleType::getInputType(
size_t idx) {
1013 const auto &portInfo = getPorts()[getPortIdForInputId(idx)];
1015 return portInfo.type;
1016 return InOutType::get(portInfo.type);
1019Type ModuleType::getOutputType(
size_t idx) {
1020 return getPorts()[getPortIdForOutputId(idx)].type;
1023SmallVector<Attribute> ModuleType::getInputNames() {
1024 SmallVector<Attribute> retval;
1025 for (
auto &p : getPorts())
1027 retval.push_back(p.name);
1031SmallVector<Attribute> ModuleType::getOutputNames() {
1032 SmallVector<Attribute> retval;
1033 for (
auto &p : getPorts())
1035 retval.push_back(p.name);
1039StringAttr ModuleType::getPortNameAttr(
size_t idx) {
1040 return getPorts()[idx].name;
1043StringRef ModuleType::getPortName(
size_t idx) {
1044 auto sa = getPortNameAttr(idx);
1046 return sa.getValue();
1050StringAttr ModuleType::getInputNameAttr(
size_t idx) {
1051 return getPorts()[getPortIdForInputId(idx)].name;
1054StringRef ModuleType::getInputName(
size_t idx) {
1055 auto sa = getInputNameAttr(idx);
1057 return sa.getValue();
1061StringAttr ModuleType::getOutputNameAttr(
size_t idx) {
1062 return getPorts()[getPortIdForOutputId(idx)].name;
1065StringRef ModuleType::getOutputName(
size_t idx) {
1066 auto sa = getOutputNameAttr(idx);
1068 return sa.getValue();
1072bool ModuleType::isOutput(
size_t idx) {
1073 auto &p = getPorts()[idx];
1074 return p.dir == ModulePort::Direction::Output;
1077FunctionType ModuleType::getFuncType() {
1078 SmallVector<Type> inputs, outputs;
1079 for (
auto p : getPorts())
1081 inputs.push_back(p.type);
1083 inputs.push_back(InOutType::get(p.type));
1085 outputs.push_back(p.type);
1086 return FunctionType::get(getContext(), inputs, outputs);
1089ArrayRef<ModulePort> ModuleType::getPorts()
const {
1090 return getImpl()->getPorts();
1093FailureOr<ModuleType> ModuleType::resolveParametricTypes(ArrayAttr parameters,
1096 SmallVector<ModulePort, 8> resolvedPorts;
1098 FailureOr<Type> resolvedType =
1100 if (failed(resolvedType))
1102 port.type = *resolvedType;
1103 resolvedPorts.push_back(port);
1105 return ModuleType::get(getContext(), resolvedPorts);
1110 case ModulePort::Direction::Input:
1112 case ModulePort::Direction::Output:
1114 case ModulePort::Direction::InOut:
1121 return ModulePort::Direction::Input;
1122 if (str ==
"output")
1123 return ModulePort::Direction::Output;
1125 return ModulePort::Direction::InOut;
1126 llvm::report_fatal_error(
"invalid direction");
1132 SmallVectorImpl<ModulePort> &ports) {
1133 return p.parseCommaSeparatedList(
1134 mlir::AsmParser::Delimiter::LessGreater, [&]() -> ParseResult {
1138 if (p.parseKeyword(&dir) || p.parseKeywordOrString(&name) ||
1139 p.parseColon() || p.parseType(type))
1142 {StringAttr::get(p.getContext(), name), type,
strToDir(dir)});
1148static void printPorts(AsmPrinter &p, ArrayRef<ModulePort> ports) {
1150 llvm::interleaveComma(ports, p, [&](
const ModulePort &port) {
1152 p.printKeywordOrString(port.
name.getValue());
1153 p <<
" : " << port.
type;
1158Type ModuleType::parse(AsmParser &odsParser) {
1159 llvm::SmallVector<ModulePort, 4> ports;
1162 return get(odsParser.getContext(), ports);
1165void ModuleType::print(AsmPrinter &odsPrinter)
const {
1170 ArrayRef<Attribute> inputNames,
1171 ArrayRef<Attribute> outputNames) {
1173 cast<FunctionType>(cast<mlir::FunctionOpInterface>(op).getFunctionType()),
1174 inputNames, outputNames);
1178 ArrayRef<Attribute> inputNames,
1179 ArrayRef<Attribute> outputNames) {
1180 SmallVector<ModulePort> ports;
1181 if (!inputNames.empty()) {
1182 for (
auto [t, n] : llvm::zip_equal(fnty.getInputs(), inputNames))
1183 if (
auto iot = dyn_cast<hw::InOutType>(t))
1184 ports.push_back({cast<StringAttr>(n), iot.getElementType(),
1185 ModulePort::Direction::InOut});
1187 ports.push_back({cast<StringAttr>(n), t, ModulePort::Direction::Input});
1189 for (
auto t : fnty.getInputs())
1190 if (auto iot = dyn_cast<
hw::InOutType>(t))
1192 {{}, iot.getElementType(), ModulePort::Direction::InOut});
1194 ports.push_back({{}, t, ModulePort::Direction::Input});
1196 if (!outputNames.empty()) {
1197 for (
auto [t, n] :
llvm::zip_equal(fnty.getResults(), outputNames))
1198 ports.push_back({cast<StringAttr>(n), t, ModulePort::Direction::Output});
1200 for (
auto t : fnty.getResults())
1201 ports.push_back({{}, t, ModulePort::Direction::Output});
1203 return ModuleType::get(fnty.getContext(), ports);
1208 size_t nextInput = 0;
1209 size_t nextOutput = 0;
1210 for (
auto [idx, p] : llvm::enumerate(
ports)) {
1229void HWDialect::registerTypes() {
1231#define GET_TYPEDEF_LIST
1232#include "circt/Dialect/HW/HWTypes.cpp.inc"
assert(baseType &&"element must be base type")
MlirType uint64_t numElements
static ModulePort::Direction strToDir(StringRef str)
static void printPorts(AsmPrinter &p, ArrayRef< ModulePort > ports)
Print out a list of named fields surrounded by <>.
static void printFields(AsmPrinter &p, ArrayRef< FieldInfo > fields)
Print out a list of named fields surrounded by <>.
static StringRef dirToStr(ModulePort::Direction dir)
static ParseResult parseHWArray(AsmParser &parser, Attribute &dim, Type &elementType)
static ParseResult parseHWElementType(AsmParser &parser, Type &elementType)
Parse and print nested HW types nicely.
static ParseResult parsePorts(AsmParser &p, SmallVectorImpl< ModulePort > &ports)
Parse a list of field names and types within <>.
static void printHWArray(AsmPrinter &printer, Attribute dim, Type elementType)
static std::pair< uint64_t, SmallVector< uint64_t > > getFieldIDsStruct(const StructType &st)
static ParseResult parseFields(AsmParser &p, SmallVectorImpl< FieldInfo > ¶meters)
Parse a list of unique field names and types within <>.
static void printHWElementType(AsmPrinter &printer, Type dim)
static unsigned getFieldID(BundleType type, unsigned index)
static unsigned getIndexForFieldID(BundleType type, unsigned fieldID)
static unsigned getMaxFieldID(FIRRTLBaseType type)
static InstancePath empty
This stores lookup tables to make manipulating and working with the IR more efficient.
mlir::Operation * getDefinition(mlir::Attribute attr) const override
Lookup a definition for 'symbol' in the cache.
Direction get(bool isOutput)
Returns an output direction if isOutput is true, otherwise returns an input direction.
Direction
The direction of a Component or Cell port.
uint64_t getWidth(Type t)
mlir::Type innerType(mlir::Type type)
std::pair< uint64_t, uint64_t > getIndexAndSubfieldID(Type type, uint64_t fieldID)
std::pair<::mlir::Type, uint64_t > getSubTypeByFieldID(Type, uint64_t fieldID)
uint64_t getMaxFieldID(Type)
llvm::hash_code hash_value(const FieldInfo &fi)
bool operator==(const FieldInfo &a, const FieldInfo &b)
ModuleType fnToMod(Operation *op, ArrayRef< Attribute > inputNames, ArrayRef< Attribute > outputNames)
bool isHWIntegerType(mlir::Type type)
Return true if the specified type is a value HW Integer type.
bool isHWValueType(mlir::Type type)
Return true if the specified type can be used as an HW value type, that is the set of types that can ...
bool isValidProbeElementType(mlir::Type type)
Return true if type is a valid probe payload.
LogicalResult aggregateAttrToAPInt(mlir::Type type, ArrayAttr attr, APInt &result)
Convert an ArrayAttr into an APInt value matching the given type.
mlir::FailureOr< mlir::Type > evaluateParametricType(mlir::Location loc, mlir::ArrayAttr parameters, mlir::Type type, bool emitErrors=true)
Returns a resolved version of 'type' wherein any parameter reference has been evaluated based on the ...
int64_t getBitWidth(mlir::Type type)
Return the hardware bit width of a type.
LogicalResult apIntToAggregateAttr(mlir::Type aggregateType, const APInt &intVal, ArrayAttr &result)
Convert an APInt value into a nested aggregate attribute matching the given HWAggregateType.
bool isHWEnumType(mlir::Type type)
Return true if the specified type is a HW Enum type.
mlir::Type getCanonicalType(mlir::Type type)
Recursively remove HW type aliases from a type and its subelements.
bool hasHWInOutType(mlir::Type type)
Return true if the specified type contains known marker types like InOutType.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
Interface for dialects to classify their types as valid probe payloads.
virtual bool isValidProbeElementType(mlir::Type type) const =0
Struct defining a field. Used in structs.
SmallVector< ModulePort > ports
The parametric data held by the storage class.
ModuleTypeStorage(ArrayRef< ModulePort > inPorts)
SmallVector< size_t > absToInput
SmallVector< size_t > outputToAbs
SmallVector< size_t > inputToAbs
SmallVector< size_t > absToOutput
Struct defining a field with an offset. Used in unions.