LLZK 3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Loading...
Searching...
No Matches
PodToScalarPass.cpp
Go to the documentation of this file.
1//===-- PodToScalarPass.cpp -------------------------------------*- C++ -*-===//
2//
3// Part of the LLZK Project, under the Apache License v2.0.
4// See LICENSE.txt for license information.
5// Copyright 2026 Project LLZK
6// SPDX-License-Identifier: Apache-2.0
7//
8//===----------------------------------------------------------------------===//
19/// POD record, replaces array-typed struct members whose element type is a POD with one parallel
20/// array member per POD record, and remembers how each original member was split for the later
21/// rewriting steps.
22///
23/// 2. Run a dialect conversion that splits arrays whose element type is a POD into parallel arrays
24/// in `llzk.nondet`, `array.*`, `constrain.eq`, `constrain.in`, `struct.readm`, `struct.writem`,
25/// `function.def`, `function.call`, and `function.return`.
26///
27/// 3. Run a dialect conversion that does the following:
28///
29/// - Replace `MemberReadOp` and `MemberWriteOp` targeting the pod-typed struct members split in
30/// step 1 so they instead perform reads and writes on the new scalar members. Reads and writes
31/// are tracked through virtual POD placeholders so the conversion can keep propagating scalar
32/// leaves instead of re-introducing aggregate POD storage.
33///
34/// - Remove optional initialization from `NewPodOp` and instead insert a list of `WritePodOp`
35/// immediately following.
36///
37/// - Split remaining direct POD values to scalars in `FuncDefOp`, `CallOp`, and `ReturnOp`.
38/// When a rewritten op still needs POD contents locally, keep them in the same virtual
39/// placeholder form for as long as possible and only materialize concrete `pod.write`
40/// operations as a fallback for unresolved uses.
41///
42/// 4. Promote pod reads and writes out of `scf.if`, `scf.for`, and `scf.while` regions when the
43/// access can be modeled as an SSA value flowing through the region boundary. This puts the
44/// pod accesses that mem2reg must eliminate into a parent block or loop-carried value.
45///
46/// 5. Run MLIR "sroa" pass to split remaining POD allocations into single-record POD allocations
47/// (to prepare for the "mem2reg" pass because its API cannot split memory by itself).
48///
49/// 6. Run MLIR "mem2reg" pass to convert all single-record POD allocations and accesses into SSA
50/// values.
51///
52/// 7. Remove POD allocations that become unread after memory promotion, then remove SSA values
53/// made dead by that cleanup.
54///
55/// Steps 5-7 are rerun while nested POD types are still being exposed, until a fixpoint.
56///
57/// Note: This transformation imposes a "last write wins" semantics on pod records. If
58/// different/configurable semantics are added in the future, some additional transformation would
59/// be necessary before/during this pass so that multiple writes to the same record can be handled
60/// properly while they still exist.
61///
62/// Note: This transformation will introduce a `nondet` op when there exists a read from a pod
63/// record that was not earlier written to.
64///
65/// Terminology: A "virtual POD" is a POD-typed placeholder whose contents are represented by SSA
66/// values for each flattened leaf record rather than by explicit `pod.write` operations on an
67/// aggregate POD object. The pass uses virtual PODs to keep propagating scalar leaves through
68/// rewrites for as long as possible because that avoids re-introducing aggregate POD storage that
69/// later stages would need to split or promote again. Concrete POD storage is materialized only
70/// when some remaining use cannot be resolved from those leaf values directly.
71///
72//===----------------------------------------------------------------------===//
73
102#include "llzk/Util/Concepts.h"
103#include "llzk/Util/TypeHelper.h"
104#include "llzk/Util/Walk.h"
105
106#include <mlir/Dialect/SCF/IR/SCF.h>
107#include <mlir/Dialect/SCF/Transforms/Patterns.h>
108#include <mlir/Pass/PassManager.h>
109#include <mlir/Transforms/DialectConversion.h>
110#include <mlir/Transforms/GreedyPatternRewriteDriver.h>
111#include <mlir/Transforms/Passes.h>
112
113#include <llvm/ADT/DenseMapInfo.h>
114#include <llvm/ADT/STLExtras.h>
115#include <llvm/ADT/TypeSwitch.h>
116#include <llvm/Support/Debug.h>
117#include <llvm/Support/raw_ostream.h>
118
119#include <functional>
120#include <limits>
121#include <optional>
122
123// Include the generated base pass class definitions.
124namespace llzk::pod {
125#define GEN_PASS_DEF_PODTOSCALARPASS
127} // namespace llzk::pod
128
129using namespace mlir;
130using namespace llzk;
131using namespace llzk::array;
132using namespace llzk::pod;
133using namespace llzk::function;
134using namespace llzk::component;
135using namespace llzk::polymorphic;
136
137#define DEBUG_TYPE "llzk-pod-to-scalar"
138
139namespace {
140
142template <typename OpTy> static OpTy preserveDiscardableAttrs(Operation *src, OpTy dst) {
143 dst->setDiscardableAttrs(src->getDiscardableAttrDictionary());
144 return dst;
145}
146
148template <typename OpTy>
149static OpTy preserveDiscardableAttrsExcept(Operation *src, OpTy dst, StringRef excludedAttr) {
150 auto original = src->getDiscardableAttrDictionary();
151 SmallVector<NamedAttribute> attrs;
152 for (NamedAttribute attr : original.getValue()) {
153 if (attr.getName().getValue() != excludedAttr) {
154 attrs.push_back(attr);
155 }
156 }
157 dst->setDiscardableAttrs(DictionaryAttr::get(src->getContext(), attrs));
158 return dst;
159}
160
162struct RecordChain {
163 SmallVector<StringAttr> nameList;
164 bool syntheticShapeCarrier = false;
165
166 RecordChain() = default;
167
168 explicit RecordChain(ArrayRef<StringAttr> names, bool syntheticShape = false)
169 : nameList(names.begin(), names.end()), syntheticShapeCarrier(syntheticShape) {}
170
172 StringRef getComponentName(size_t componentIndex) const {
173 assert(componentIndex < nameList.size() && "component index must be in range");
174 bool syntheticComponent = syntheticShapeCarrier && componentIndex + 1 == nameList.size();
175 return syntheticComponent ? StringRef("shape") : nameList[componentIndex].getValue();
176 }
177
179 RecordChain withPrefix(ArrayRef<StringAttr> prefix) const {
180 SmallVector<StringAttr> fullChain(prefix.begin(), prefix.end());
181 llvm::append_range(fullChain, nameList);
182 return RecordChain(fullChain, syntheticShapeCarrier);
183 }
184
186 StringAttr getFlattenedMemberName(MLIRContext *ctx, StringAttr memberName) const {
187 std::string flatName;
188 llvm::raw_string_ostream os(flatName);
189 os << memberName.getValue();
190 for (size_t i = 0; i < nameList.size(); ++i) {
191 os << '_' << getComponentName(i);
192 }
193 return StringAttr::get(ctx, flatName);
194 }
195
196 bool operator==(const RecordChain &other) const {
197 return syntheticShapeCarrier == other.syntheticShapeCarrier && nameList == other.nameList;
198 }
199};
200
202struct CompatiblePodLeafMaterializationKey {
203 Value source;
204 Type podType;
205
206 bool operator==(const CompatiblePodLeafMaterializationKey &other) const {
207 return source == other.source && podType == other.podType;
208 }
209};
210
211using CompatiblePodLeafMaterializationMap =
212 DenseMap<CompatiblePodLeafMaterializationKey, SmallVector<Value>>;
213
214} // namespace
215
216namespace llvm {
217
218template <> struct DenseMapInfo<RecordChain> {
219 static RecordChain getEmptyKey() {
220 return RecordChain {{DenseMapInfo<StringAttr>::getEmptyKey()}};
221 }
222
223 static RecordChain getTombstoneKey() {
224 return RecordChain {{DenseMapInfo<StringAttr>::getTombstoneKey()}};
225 }
226
227 static unsigned getHashValue(const RecordChain &chain) {
228 return llvm::hash_combine(
229 llvm::hash_combine_range(chain.nameList.begin(), chain.nameList.end()),
230 chain.syntheticShapeCarrier
231 );
232 }
233
234 static bool isEqual(const RecordChain &lhs, const RecordChain &rhs) { return lhs == rhs; }
235};
236
237template <> struct DenseMapInfo<CompatiblePodLeafMaterializationKey> {
238 static CompatiblePodLeafMaterializationKey getEmptyKey() {
239 return {DenseMapInfo<Value>::getEmptyKey(), DenseMapInfo<Type>::getEmptyKey()};
240 }
241
242 static CompatiblePodLeafMaterializationKey getTombstoneKey() {
243 return {DenseMapInfo<Value>::getTombstoneKey(), DenseMapInfo<Type>::getTombstoneKey()};
244 }
245
246 static unsigned getHashValue(const CompatiblePodLeafMaterializationKey &key) {
247 return llvm::hash_combine(key.source, key.podType);
248 }
249
250 static bool isEqual(
251 const CompatiblePodLeafMaterializationKey &lhs, const CompatiblePodLeafMaterializationKey &rhs
252 ) {
253 return lhs == rhs;
254 }
255};
256
257} // namespace llvm
258
259namespace {
260
262inline static bool isSamePodRecord(ReadPodOp readOp, Value podRef, StringAttr recordName) {
263 return readOp.getPodRef() == podRef && readOp.getRecordNameAttr() == recordName;
264}
265
267inline static bool isSamePodRecord(WritePodOp writeOp, Value podRef, StringAttr recordName) {
268 return writeOp.getPodRef() == podRef && writeOp.getRecordNameAttr() == recordName;
269}
270
272static bool hasNestedWriteToRecord(Operation &op, Value podRef, StringAttr recordName) {
273 return walkContainsMatch<WritePodOp>(op, [&](WritePodOp writeOp) {
274 return writeOp.getOperation() != &op && isSamePodRecord(writeOp, podRef, recordName);
275 });
276}
277
279static bool hasNestedWriteToPod(Operation &op, Value podRef) {
280 return walkContainsMatch<WritePodOp>(op, [&](WritePodOp writeOp) {
281 return writeOp.getOperation() != &op && writeOp.getPodRef() == podRef;
282 });
283}
284
286static bool hasReadFromRecord(Operation &op, Value podRef, StringAttr recordName) {
287 return walkContainsMatch<ReadPodOp>(op, [&podRef, &recordName](ReadPodOp readOp) {
288 return isSamePodRecord(readOp, podRef, recordName);
289 });
290}
291
293static bool hasValueUse(Operation &op, Value value) {
294 return walkContainsMatch<Operation *>(op, [&value](Operation *nestedOp) {
295 return llvm::is_contained(nestedOp->getOperands(), value);
296 });
297}
298
304static WritePodOp
305findNearestForwardableWriteBefore(Operation *cursor, Value podRef, StringAttr recordName) {
306 for (Operation *op = cursor; op; op = op->getPrevNode()) {
307 if (!hasValueUse(*op, podRef)) {
308 continue;
309 }
310 auto writeOp = dyn_cast<WritePodOp>(op);
311 return writeOp && isSamePodRecord(writeOp, podRef, recordName) ? writeOp : nullptr;
312 }
313 return nullptr;
314}
315
317inline static WritePodOp findNearestForwardableWriteInBlock(ReadPodOp readOp) {
318 return findNearestForwardableWriteBefore(
319 readOp->getPrevNode(), readOp.getPodRef(), readOp.getRecordNameAttr()
320 );
321}
322
329static WritePodOp findNearestForwardableWrite(ReadPodOp readOp) {
330 Value podRef = readOp.getPodRef();
331 StringAttr recordName = readOp.getRecordNameAttr();
332 Operation *loopBoundary = llzk::pod::detail::findNearestLoopCarriedPodAccess(readOp);
333
334 for (Operation *cursor = readOp.getOperation(); cursor && cursor->getBlock();) {
335 if (WritePodOp writeOp =
336 findNearestForwardableWriteBefore(cursor->getPrevNode(), podRef, recordName)) {
337 return writeOp;
338 }
339
340 Operation *parentOp = cursor->getBlock()->getParentOp();
341 if (parentOp == loopBoundary) {
342 break;
343 }
344 cursor = parentOp;
345 }
346
347 return nullptr;
348}
349
351inline static PodType splittablePod(PodType pt) { return pt; }
352
354inline static PodType splittablePod(Type t) {
355 if (PodType pt = dyn_cast<PodType>(t)) {
356 return splittablePod(pt);
357 } else {
358 return nullptr;
359 }
360}
361
364inline static bool containsSplittablePodType(ArrayRef<Type> types) {
365 for (Type t : types) {
366 if (splittablePod(t)) {
367 return true;
368 }
369 }
370 return false;
371}
372
375template <typename T> inline static bool containsSplittablePodType(ValueTypeRange<T> types) {
376 for (Type t : types) {
377 if (splittablePod(t)) {
378 return true;
379 }
380 }
381 return false;
382}
383
385inline static ArrayType splittablePodArray(ArrayType at) {
386 return isa<PodType>(at.getElementType()) ? at : nullptr;
387}
388
390inline static ArrayType splittablePodArray(Type t) {
391 if (ArrayType at = dyn_cast<ArrayType>(t)) {
392 return splittablePodArray(at);
393 }
394 return nullptr;
395}
396
398inline static bool containsSplittablePodArrayType(ArrayRef<Type> types) {
399 return llvm::any_of(types, [](Type t) { return splittablePodArray(t); });
400}
401
403template <typename T> inline static bool containsSplittablePodArrayType(ValueTypeRange<T> types) {
404 return llvm::any_of(types, [](Type t) { return splittablePodArray(t); });
405}
406
408inline static StringAttr getPodArrayShapeCarrierMarker(MLIRContext *ctx) {
409 return StringAttr::get(ctx, "__llzk_pod_array_shape");
410}
411
413[[maybe_unused]] inline static bool isPodArrayShapeCarrierMarker(StringAttr recordName) {
414 return recordName && recordName == getPodArrayShapeCarrierMarker(recordName.getContext());
415}
416
418inline constexpr llvm::StringLiteral DEFERRED_POD_ARRAY_LENGTH_ATTR =
419 "llzk.deferred_pod_array_length";
420
422inline constexpr llvm::StringLiteral RAGGED_DYNAMIC_NESTED_LEAF_ATTR =
423 "llzk.ragged_nested_dynamic_leaf";
424
426inline constexpr llvm::StringLiteral RAGGED_AFFINE_NESTED_LEAF_ATTR =
427 "llzk.ragged_nested_affine_leaf";
428
430inline static ArrayType getPodArrayShapeCarrierType(ArrayType arrTy) {
431 return arrTy.cloneWith(NoneType::get(arrTy.getContext()));
432}
433
435inline static bool hasWildcardArrayDimensions(ArrayRef<Attribute> dims) {
436 return llvm::any_of(dims, [](Attribute dim) {
437 if (auto intAttr = llvm::dyn_cast<IntegerAttr>(dim)) {
438 return isDynamic(intAttr);
439 }
440 return false;
441 });
442}
443
445inline static bool hasAffineArrayDimensions(ArrayRef<Attribute> dims) {
446 return llvm::any_of(dims, [](Attribute dim) { return llvm::isa<AffineMapAttr>(dim); });
447}
448
449// pre-declared (see definition below)
450template <typename Fn>
451static void forEachPodLeaf(PodType podTy, SmallVectorImpl<StringAttr> &recordChain, Fn &&callback);
452
454static size_t splitPodArrayTypeTo(
455 Type t, SmallVectorImpl<Type> &collect, SmallVector<RecordChain> *splitIds = nullptr
456) {
457 if (ArrayType at = splittablePodArray(t)) {
458 auto podTy = llvm::cast<PodType>(at.getElementType());
459 SmallVector<StringAttr> recordChain;
460 size_t originalSize = collect.size();
461 forEachPodLeaf(podTy, recordChain, [&](RecordChain id, Type leafType) {
462 collect.push_back(flattenArrayElementType(at, leafType));
463 if (splitIds) {
464 splitIds->push_back(std::move(id));
465 }
466 });
467 return collect.size() - originalSize;
468 }
469
470 collect.push_back(t);
471 return 1;
472}
473
475template <typename TypeCollection>
476void splitPodArrayTypeTo(
477 TypeCollection types, SmallVectorImpl<Type> &collect, SmallVector<size_t> *originalIdxToSize
478) {
479 for (Type t : types) {
480 size_t count = splitPodArrayTypeTo(t, collect);
481 if (originalIdxToSize) {
482 originalIdxToSize->push_back(count);
483 }
484 }
485}
486
488static bool hasZeroLeafPodArraySplit(ArrayType arrTy) {
489 SmallVector<Type> splitTypes;
490 splitPodArrayTypeTo(arrTy, splitTypes);
491 return splitTypes.empty();
492}
493
498static size_t countRankPreservingSplitPodArrayLeaves(ArrayType arrTy) {
499 SmallVector<Type> splitTypes;
500 splitPodArrayTypeTo(arrTy, splitTypes);
501 size_t originalRank = arrTy.getDimensionSizes().size();
502 return llvm::count_if(splitTypes, [originalRank](Type splitType) {
503 auto splitArrTy = llvm::dyn_cast<ArrayType>(splitType);
504 return splitArrTy && splitArrTy.getDimensionSizes().size() == originalRank;
505 });
506}
507
514static bool needsPodArrayShapeCarrier(ArrayType arrTy) {
515 if (hasZeroLeafPodArraySplit(arrTy)) {
516 return true;
517 }
518 if (!hasWildcardArrayDimensions(arrTy.getDimensionSizes()) &&
519 !hasAffineArrayDimensions(arrTy.getDimensionSizes())) {
520 return false;
521 }
522 return countRankPreservingSplitPodArrayLeaves(arrTy) != 1;
523}
524
526static void collectConvertedPodArrayRecordInfos(
527 ArrayType arrTy, SmallVector<RecordChain> &splitIds, SmallVectorImpl<Type> &splitTypes
528) {
529 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
530 if (needsPodArrayShapeCarrier(arrTy)) {
531 splitIds.push_back(RecordChain({getPodArrayShapeCarrierMarker(arrTy.getContext())}, true));
532 splitTypes.push_back(getPodArrayShapeCarrierType(arrTy));
533 }
534}
535
537static StringRef getRaggedNestedLeafAttrName(ArrayType arrTy, Type splitType) {
538 auto splitArrTy = llvm::dyn_cast<ArrayType>(splitType);
539 if (!splitArrTy) {
540 return {};
541 }
542
543 size_t originalRank = arrTy.getDimensionSizes().size();
544 if (splitArrTy.getDimensionSizes().size() <= originalRank) {
545 return {};
546 }
547
548 ArrayRef<Attribute> nestedDims = splitArrTy.getDimensionSizes().drop_front(originalRank);
549 if (hasWildcardArrayDimensions(nestedDims)) {
550 return RAGGED_DYNAMIC_NESTED_LEAF_ATTR;
551 }
552 if (hasAffineArrayDimensions(nestedDims)) {
553 return RAGGED_AFFINE_NESTED_LEAF_ATTR;
554 }
555 return {};
556}
557
559static StringRef getRaggedNestedLeafKind(ArrayType arrTy, Type splitType) {
560 StringRef attrName = getRaggedNestedLeafAttrName(arrTy, splitType);
561 if (attrName == RAGGED_DYNAMIC_NESTED_LEAF_ATTR) {
562 return "dynamic";
563 }
564 if (attrName == RAGGED_AFFINE_NESTED_LEAF_ATTR) {
565 return "affine";
566 }
567 return {};
568}
569
571static StringRef getRaggedNestedLeafKind(ArrayType arrTy) {
572 SmallVector<Type> splitTypes;
573 splitPodArrayTypeTo(arrTy, splitTypes);
574 for (Type splitType : splitTypes) {
575 if (StringRef raggedKind = getRaggedNestedLeafKind(arrTy, splitType); !raggedKind.empty()) {
576 return raggedKind;
577 }
578 }
579 return {};
580}
581
583static Value
584tagRaggedNestedLeafValue(OpBuilder &bldr, Location loc, Value value, StringRef attrName) {
585 if (attrName.empty()) {
586 return value;
587 }
588 auto cast = bldr.create<UnrealizedConversionCastOp>(loc, TypeRange {value.getType()}, value);
589 cast->setAttr(attrName, UnitAttr::get(bldr.getContext()));
590 return cast.getResult(0);
591}
592
594static StringRef getTaggedRaggedNestedLeafKind(Value value) {
595 while (true) {
596 if (auto cast = value.getDefiningOp<UnifiableCastOp>()) {
597 value = cast.getInput();
598 continue;
599 }
600 auto cast = value.getDefiningOp<UnrealizedConversionCastOp>();
601 if (!cast) {
602 break;
603 }
604 if (cast->hasAttr(RAGGED_DYNAMIC_NESTED_LEAF_ATTR)) {
605 return "dynamic";
606 }
607 if (cast->hasAttr(RAGGED_AFFINE_NESTED_LEAF_ATTR)) {
608 return "affine";
609 }
610 if (cast->getNumOperands() != 1) {
611 break;
612 }
613 value = cast.getOperand(0);
614 }
615 return {};
616}
617
619static LogicalResult
620rejectRaggedNestedLeafEquality(constrain::EmitEqualityOp op, Value lhsLeaf, Value rhsLeaf) {
621 StringRef raggedKind = getTaggedRaggedNestedLeafKind(lhsLeaf);
622 if (raggedKind.empty()) {
623 raggedKind = getTaggedRaggedNestedLeafKind(rhsLeaf);
624 }
625 if (raggedKind.empty()) {
626 return success();
627 }
628 return op.emitOpError() << "cannot lower nested " << raggedKind
629 << " array leaf equality after reading an array-of-POD element without "
630 "per-element shape witnesses";
631}
632
634static LogicalResult rejectRaggedNestedLeafContainment(constrain::EmitContainmentOp op) {
635 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getLhs());
636 if (raggedKind.empty()) {
637 raggedKind = getTaggedRaggedNestedLeafKind(op.getRhs());
638 }
639 if (raggedKind.empty()) {
640 return success();
641 }
642 return op.emitOpError() << "cannot lower nested " << raggedKind
643 << " array leaf containment after reading an array-of-POD element "
644 "without per-element shape witnesses";
645}
646
648template <typename ValueRangeLike>
649static LogicalResult
650rejectRaggedNestedLeafBoundaryCrossing(Operation *op, ValueRangeLike values, StringRef boundary) {
651 StringRef raggedKind;
652 for (Value value : values) {
653 raggedKind = getTaggedRaggedNestedLeafKind(value);
654 if (!raggedKind.empty()) {
655 break;
656 }
657 }
658 if (raggedKind.empty()) {
659 return success();
660 }
661 return op->emitOpError() << "cannot pass nested " << raggedKind << " array leaf " << boundary
662 << " after reading an array-of-POD element without per-element shape "
663 "witnesses";
664}
665
667static Value peelUnifiableCasts(Value value) {
668 while (auto cast = value.getDefiningOp<UnifiableCastOp>()) {
669 value = cast.getInput();
670 }
671 return value;
672}
673
675static Type getFlattenedTypeAlongPath(
676 Type type, ArrayRef<StringAttr> recordChain, bool syntheticShapeCarrier = false
677) {
678 if (recordChain.empty()) {
679 return type;
680 }
681
682 if (PodType podTy = dyn_cast<PodType>(type)) {
683 Type nextType = podTy.getRecordMap().lookup(recordChain.front().getValue());
684 assert(nextType && "record path must exist in the containing POD");
685 return getFlattenedTypeAlongPath(nextType, recordChain.drop_front(), syntheticShapeCarrier);
686 }
687
688 if (ArrayType arrTy = splittablePodArray(type)) {
689 if (syntheticShapeCarrier && recordChain.size() == 1) {
690 assert(
691 isPodArrayShapeCarrierMarker(recordChain.front()) &&
692 "synthetic shape carrier must use the reserved shape marker"
693 );
694 assert(needsPodArrayShapeCarrier(arrTy) && "shape marker requires an explicit array carrier");
695 return getPodArrayShapeCarrierType(arrTy);
696 }
697
698 auto elemPodTy = llvm::cast<PodType>(arrTy.getElementType());
699 Type nextType = elemPodTy.getRecordMap().lookup(recordChain.front().getValue());
700 assert(nextType && "record path must exist in the POD array element type");
702 arrTy, getFlattenedTypeAlongPath(nextType, recordChain.drop_front(), syntheticShapeCarrier)
703 );
704 }
705
706 llvm_unreachable("record path cannot continue through a non-POD leaf");
707}
708
710inline static Type getFlattenedTypeAlongPath(Type type, const RecordChain &recordChain) {
711 return getFlattenedTypeAlongPath(
712 type, ArrayRef(recordChain.nameList), recordChain.syntheticShapeCarrier
713 );
714}
715
717template <typename Fn>
718static void forEachPodLeaf(PodType podTy, SmallVectorImpl<StringAttr> &recordChain, Fn &&callback) {
719 std::function<void(Type)> walk = [&](Type type) {
720 if (PodType nestedPodTy = llvm::dyn_cast<PodType>(type)) {
721 for (RecordAttr record : nestedPodTy.getRecords()) {
722 recordChain.push_back(record.getName());
723 walk(record.getType());
724 recordChain.pop_back();
725 }
726 } else if (ArrayType arrTy = splittablePodArray(type)) {
727 auto elemPodTy = llvm::cast<PodType>(arrTy.getElementType());
728 for (RecordAttr record : elemPodTy.getRecords()) {
729 recordChain.push_back(record.getName());
730 walk(flattenArrayElementType(arrTy, record.getType()));
731 recordChain.pop_back();
732 }
733 if (needsPodArrayShapeCarrier(arrTy)) {
734 recordChain.push_back(getPodArrayShapeCarrierMarker(arrTy.getContext()));
735 callback(RecordChain(recordChain, true), getPodArrayShapeCarrierType(arrTy));
736 recordChain.pop_back();
737 }
738 } else {
739 callback(RecordChain(recordChain), type);
740 }
741 };
742
743 walk(podTy);
744}
745
748size_t splitPodTypeTo(Type t, SmallVector<Type> &collect) {
749 if (PodType pt = splittablePod(t)) {
750 SmallVector<StringAttr> recordChain;
751 size_t originalSize = collect.size();
752 forEachPodLeaf(pt, recordChain, [&collect](const RecordChain &, Type leafType) {
753 collect.push_back(leafType);
754 });
755 return collect.size() - originalSize;
756 } else {
757 collect.push_back(t);
758 return 1;
759 }
760}
761
763template <typename TypeCollection>
764inline void splitPodTypeTo(
765 TypeCollection types, SmallVector<Type> &collect, SmallVector<size_t> *originalIdxToSize
766) {
767 for (Type t : types) {
768 size_t count = splitPodTypeTo(t, collect);
769 if (originalIdxToSize) {
770 originalIdxToSize->push_back(count);
771 }
772 }
773}
774
777template <typename TypeCollection>
778inline SmallVector<Type>
779splitPodType(TypeCollection types, SmallVector<size_t> *originalIdxToSize = nullptr) {
780 SmallVector<Type> collect;
781 splitPodTypeTo(types, collect, originalIdxToSize);
782 return collect;
783}
784
786static Value
787getConvertedPodArrayShapeCarrierIfPresent(ArrayType arrTy, ValueRange convertedValues) {
788 SmallVector<Type> splitTypes;
789 splitPodArrayTypeTo(arrTy, splitTypes);
790 if (!needsPodArrayShapeCarrier(arrTy) || convertedValues.size() != splitTypes.size() + 1) {
791 return {};
792 }
793 return convertedValues.back();
794}
795
802static size_t convertPodArrayTypeTo(Type t, SmallVectorImpl<Type> &collect) {
803 if (ArrayType arrTy = splittablePodArray(t)) {
804 size_t oldSize = collect.size();
805 splitPodArrayTypeTo(arrTy, collect);
806 if (needsPodArrayShapeCarrier(arrTy)) {
807 collect.push_back(getPodArrayShapeCarrierType(arrTy));
808 }
809 return collect.size() - oldSize;
810 }
811
812 collect.push_back(t);
813 return 1;
814}
815
817template <typename TypeCollection>
818inline void convertPodArrayTypesTo(
819 TypeCollection types, SmallVectorImpl<Type> &collect,
820 SmallVector<size_t> *originalIdxToSize = nullptr
821) {
822 if (originalIdxToSize) {
823 originalIdxToSize->reserve(types.size());
824 }
825 for (Type t : types) {
826 size_t count = convertPodArrayTypeTo(t, collect);
827 if (originalIdxToSize) {
828 originalIdxToSize->push_back(count);
829 }
830 }
831}
832
834template <typename TypeCollection>
835inline static SmallVector<Type>
836convertPodArrayTypes(TypeCollection types, SmallVector<size_t> *originalIdxToSize = nullptr) {
837 SmallVector<Type> collect;
838 convertPodArrayTypesTo(types, collect, originalIdxToSize);
839 return collect;
840}
841
843static SmallVector<std::string> getSplitPodArrayRecordNameSuffixes(Type type) {
844 SmallVector<std::string> suffixes;
845 if (ArrayType at = splittablePodArray(type)) {
846 SmallVector<RecordChain> splitIds;
847 SmallVector<Type> ignoredTypes;
848 splitPodArrayTypeTo(at, ignoredTypes, &splitIds);
849 suffixes.reserve(splitIds.size());
850 for (const RecordChain &id : splitIds) {
851 std::string suffix;
852 llvm::raw_string_ostream os(suffix);
853 for (size_t i = 0; i < id.nameList.size(); ++i) {
854 os << '.' << id.getComponentName(i);
855 }
856 suffixes.push_back(std::move(suffix));
857 }
858 if (needsPodArrayShapeCarrier(at)) {
859 suffixes.push_back(".shape");
860 }
861 }
862 return suffixes;
863}
864
870static Value castValueToTypeIfNeeded(OpBuilder &bldr, Location loc, Value value, Type targetType) {
871 if (value.getType() == targetType) {
872 return value;
873 }
874 assert(typesUnify(value.getType(), targetType) && "expected compatible rewritten types");
875 return bldr.create<UnifiableCastOp>(loc, targetType, value);
876}
877
879inline static ReadPodOp
880genRead(OpBuilder &bldr, Location loc, Value podRef, StringAttr recordName) {
881 Type resultType =
882 llvm::cast<PodType>(podRef.getType()).getRecordMap().lookup(recordName.getValue());
883 return bldr.create<ReadPodOp>(loc, resultType, podRef, recordName);
884}
885
887inline static WritePodOp
888genWrite(OpBuilder &bldr, Location loc, Value podRef, StringAttr recordName, Value value) {
889 Type recordType =
890 llvm::cast<PodType>(podRef.getType()).getRecordMap().lookup(recordName.getValue());
891 return bldr.create<WritePodOp>(
892 loc, podRef, recordName, castValueToTypeIfNeeded(bldr, loc, value, recordType)
893 );
894}
895
897inline static Value getSingleConvertedValue(ValueRange values) {
898 assert(values.size() == 1 && "expected a 1:1 converted value range");
899 return values.front();
900}
901
903inline static size_t getSplitPodArrayLeafCount(ArrayType arrTy) {
904 SmallVector<Type> splitTypes;
905 splitPodArrayTypeTo(arrTy, splitTypes);
906 return splitTypes.size();
907}
908
910static ValueRange getConvertedPodArrayLeafValues(ArrayType arrTy, ValueRange convertedValues) {
911 size_t leafCount = getSplitPodArrayLeafCount(arrTy);
912 return convertedValues.take_front(
913 convertedValues.size() < leafCount ? convertedValues.size() : leafCount
914 );
915}
916
918static bool hasEarlierWriteToRecordInBlock(Operation *op, Value podRef, StringAttr recordName) {
919 for (Operation &candidate : *op->getBlock()) {
920 if (&candidate == op) {
921 return false;
922 }
923 if (auto writeOp = dyn_cast<WritePodOp>(&candidate)) {
924 if (isSamePodRecord(writeOp, podRef, recordName)) {
925 return true;
926 }
927 } else if (hasNestedWriteToRecord(candidate, podRef, recordName)) {
928 return true;
929 }
930 }
931 return false;
932}
933
936static bool hasEarlierWriteToRecord(Operation *op, Value podRef, StringAttr recordName) {
937 for (Operation *cursor = op; cursor && cursor->getBlock();) {
938 if (hasEarlierWriteToRecordInBlock(cursor, podRef, recordName)) {
939 return true;
940 }
941 cursor = cursor->getBlock()->getParentOp();
942 }
943 return false;
944}
945
947inline static bool hasEarlierWriteInBlock(ReadPodOp readOp) {
948 return hasEarlierWriteToRecordInBlock(
949 readOp.getOperation(), readOp.getPodRef(), readOp.getRecordNameAttr()
950 );
951}
952
955inline static bool hasEarlierWrite(ReadPodOp readOp) {
956 return hasEarlierWriteToRecord(
957 readOp.getOperation(), readOp.getPodRef(), readOp.getRecordNameAttr()
958 );
959}
960
962static bool isFreshUnwrittenPodRead(ReadPodOp readOp) {
963 NewPodOp newPod = readOp.getPodRef().getDefiningOp<NewPodOp>();
964 if (!newPod) {
965 return false;
966 }
967 auto isReadOpRecordName = [&readOp](Attribute attr) {
968 return attr == readOp.getRecordNameAttr();
969 };
970 return llvm::none_of(newPod.getInitializedRecords(), isReadOpRecordName) &&
971 !hasEarlierWrite(readOp);
972}
973
975static void collectNewPodMapAttrs(Type type, SmallVectorImpl<AffineMapAttr> &mapAttrs) {
976 llvm::TypeSwitch<Type, void>(type)
977 .Case([&mapAttrs](PodType podTy) {
978 for (RecordAttr record : podTy.getRecords()) {
979 collectNewPodMapAttrs(record.getType(), mapAttrs);
980 }
981 })
982 .Case([&mapAttrs](ArrayType arrTy) {
983 for (Attribute dimSize : arrTy.getDimensionSizes()) {
984 if (auto mapAttr = llvm::dyn_cast<AffineMapAttr>(dimSize)) {
985 mapAttrs.push_back(mapAttr);
986 }
987 }
988 })
989 .Case([&mapAttrs](component::StructType structTy) {
990 if (ArrayAttr params = structTy.getParams()) {
991 for (Attribute param : params) {
992 if (auto mapAttr = llvm::dyn_cast<AffineMapAttr>(param)) {
993 mapAttrs.push_back(mapAttr);
994 }
995 }
996 }
997 }).Default([](Type) {});
998}
999
1004struct ArrayInstantiationInfo {
1005 SmallVector<SmallVector<Value>> mapOperandStorage;
1006 SmallVector<int32_t> numDimsPerMap;
1007};
1008
1010static std::optional<ArrayInstantiationInfo>
1011tryGetFreshUnwrittenPodReadInstantiationInfo(ReadPodOp readOp) {
1012 if (!isFreshUnwrittenPodRead(readOp)) {
1013 return std::nullopt;
1014 }
1015
1016 NewPodOp newPod = readOp.getPodRef().getDefiningOp<NewPodOp>();
1017 if (!newPod || newPod.getMapOperands().empty()) {
1018 return std::nullopt;
1019 }
1020
1021 size_t groupBegin = 0;
1022 ArrayRef<int32_t> allNumDims = newPod.getNumDimsPerMap();
1023 PodType newPodTy = newPod.getType(); // variable to avoid "dangling-reference"
1024 for (RecordAttr record : newPodTy.getRecords()) {
1025 SmallVector<AffineMapAttr> recordMapAttrs;
1026 collectNewPodMapAttrs(record.getType(), recordMapAttrs);
1027 size_t recordGroupCount = recordMapAttrs.size();
1028 if (record.getName() == readOp.getRecordNameAttr()) {
1029 if (recordGroupCount == 0 || newPod.getMapOperands().size() < groupBegin + recordGroupCount) {
1030 return std::nullopt;
1031 }
1032
1033 ArrayInstantiationInfo info;
1034 info.mapOperandStorage.reserve(recordGroupCount);
1035 info.numDimsPerMap.reserve(recordGroupCount);
1036 for (size_t i = 0; i < recordGroupCount; ++i) {
1037 OperandRange group = newPod.getMapOperands()[groupBegin + i];
1038 info.mapOperandStorage.emplace_back(group.begin(), group.end());
1039 int32_t numDims =
1040 groupBegin + i < allNumDims.size()
1041 ? allNumDims[groupBegin + i]
1042 : llzk::checkedCast<int32_t>(recordMapAttrs[i].getAffineMap().getNumDims());
1043 info.numDimsPerMap.push_back(numDims);
1044 }
1045 return info;
1046 }
1047 groupBegin += recordGroupCount;
1048 }
1049
1050 return std::nullopt;
1051}
1052
1054static size_t countRecursiveAffineMapAttrs(Type type) {
1055 return llvm::TypeSwitch<Type, size_t>(type)
1056 .Case([](PodType podTy) {
1057 size_t count = 0;
1058 for (RecordAttr record : podTy.getRecords()) {
1059 count += countRecursiveAffineMapAttrs(record.getType());
1060 }
1061 return count;
1062 })
1063 .Case([](ArrayType arrTy) {
1064 size_t count = 0;
1065 for (Attribute dimSize : arrTy.getDimensionSizes()) {
1066 if (llvm::isa<AffineMapAttr>(dimSize)) {
1067 ++count;
1068 }
1069 }
1070 return count + countRecursiveAffineMapAttrs(arrTy.getElementType());
1071 })
1072 .Case([](StructType structTy) {
1073 size_t count = 0;
1074 if (ArrayAttr params = structTy.getParams()) {
1075 for (Attribute param : params) {
1076 if (llvm::isa<AffineMapAttr>(param)) {
1077 ++count;
1078 } else if (auto typeAttr = llvm::dyn_cast<TypeAttr>(param)) {
1079 count += countRecursiveAffineMapAttrs(typeAttr.getValue());
1080 }
1081 }
1082 }
1083 return count;
1084 }).Default([](Type) { return 0; });
1085}
1086
1092static std::optional<ArrayInstantiationInfo> tryGetArrayInstantiationInfo(Value value) {
1093 while (auto cast = value.getDefiningOp<UnifiableCastOp>()) {
1094 value = cast.getInput();
1095 }
1096
1097 if (ReadPodOp read = value.getDefiningOp<ReadPodOp>()) {
1098 if (WritePodOp write = findNearestForwardableWrite(read)) {
1099 return tryGetArrayInstantiationInfo(write.getValue());
1100 }
1101 return tryGetFreshUnwrittenPodReadInstantiationInfo(read);
1102 }
1103
1104 auto create = value.getDefiningOp<CreateArrayOp>();
1105 if (!create) {
1106 return std::nullopt;
1107 }
1108
1109 ArrayInstantiationInfo info;
1110 info.mapOperandStorage.reserve(create.getMapOperands().size());
1111 for (OperandRange group : create.getMapOperands()) {
1112 info.mapOperandStorage.emplace_back(group.begin(), group.end());
1113 }
1114
1115 if (DenseI32ArrayAttr numDimsPerMap = create.getNumDimsPerMapAttr()) {
1116 llvm::append_range(info.numDimsPerMap, numDimsPerMap.asArrayRef());
1117 }
1118
1119 return info;
1120}
1121
1126static Value materializeArrayLengthCarrier(
1127 Value originalArrRef, ArrayType originalArrTy, Location loc, OpBuilder &rewriter
1128) {
1129 ArrayType carrierTy = getPodArrayShapeCarrierType(originalArrTy);
1130
1131 if (auto create = originalArrRef.getDefiningOp<CreateArrayOp>()) {
1132 if (create.getMapOperands().empty()) {
1133 return rewriter.create<CreateArrayOp>(loc, carrierTy);
1134 }
1135
1136 SmallVector<ValueRange> mapOperands;
1137 mapOperands.reserve(create.getMapOperands().size());
1138 for (OperandRange mapOperandGroup : create.getMapOperands()) {
1139 mapOperands.push_back(mapOperandGroup);
1140 }
1141 return rewriter.create<CreateArrayOp>(
1142 loc, carrierTy, mapOperands, create.getNumDimsPerMapAttr()
1143 );
1144 }
1145
1146 if (std::optional<ArrayInstantiationInfo> instantiation =
1147 tryGetArrayInstantiationInfo(originalArrRef)) {
1148 if (instantiation->mapOperandStorage.empty()) {
1149 return rewriter.create<CreateArrayOp>(loc, carrierTy);
1150 }
1151
1152 SmallVector<ValueRange> mapOperands;
1153 mapOperands.reserve(instantiation->mapOperandStorage.size());
1154 for (const SmallVector<Value> &group : instantiation->mapOperandStorage) {
1155 mapOperands.push_back(group);
1156 }
1157 return rewriter.create<CreateArrayOp>(
1158 loc, carrierTy, mapOperands, ArrayRef<int32_t>(instantiation->numDimsPerMap)
1159 );
1160 }
1161
1162 bool hasAffineDims = llvm::any_of(originalArrTy.getDimensionSizes(), [](Attribute dimSize) {
1163 return llvm::isa<AffineMapAttr>(dimSize);
1164 });
1165 if (!hasAffineDims) {
1166 return rewriter.create<CreateArrayOp>(loc, carrierTy);
1167 }
1168
1169 return rewriter.create<NonDetOp>(loc, carrierTy);
1170}
1171
1177static Value materializeExtractedPodArrayShapeCarrier(
1178 ExtractArrayOp op, ArrayType resultTy, Value originalArrRef, ValueRange convertedArrRefs,
1179 ArrayRef<Value> indices, ConversionPatternRewriter &rewriter
1180) {
1181 ArrayType originalArrTy = llvm::cast<ArrayType>(originalArrRef.getType());
1182 Value sourceCarrier = getConvertedPodArrayShapeCarrierIfPresent(originalArrTy, convertedArrRefs);
1183 if (!sourceCarrier) {
1184 sourceCarrier =
1185 materializeArrayLengthCarrier(originalArrRef, originalArrTy, op.getLoc(), rewriter);
1186 }
1187
1188 return rewriter.create<ExtractArrayOp>(
1189 op.getLoc(), getPodArrayShapeCarrierType(resultTy), sourceCarrier, indices
1190 );
1191}
1192
1194static Value getRankPreservingConvertedPodArrayLeaf(ArrayType arrTy, ValueRange convertedValues) {
1195 size_t originalRank = arrTy.getDimensionSizes().size();
1196 for (Value arrRef : getConvertedPodArrayLeafValues(arrTy, convertedValues)) {
1197 auto leafArrTy = llvm::dyn_cast<ArrayType>(arrRef.getType());
1198 if (leafArrTy && leafArrTy.getDimensionSizes().size() == originalRank) {
1199 return arrRef;
1200 }
1201 }
1202 return {};
1203}
1204
1206static Value getConvertedPodArrayShapeSource(ArrayType arrTy, ValueRange convertedValues) {
1207 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
1208 return carrier;
1209 }
1210 return getRankPreservingConvertedPodArrayLeaf(arrTy, convertedValues);
1211}
1212
1214template <typename RangeOfRanges>
1215inline static SmallVector<Value> flattenConvertedValues(RangeOfRanges ranges) {
1216 SmallVector<Value> values;
1217 for (ValueRange range : ranges) {
1218 llvm::append_range(values, range);
1219 }
1220 return values;
1221}
1222
1224struct FlattenedConvertedValueRangeStorage {
1225 SmallVector<SmallVector<Value>> storage;
1226 SmallVector<ValueRange> ranges;
1227
1228 template <typename RangeOfRanges>
1229 explicit FlattenedConvertedValueRangeStorage(RangeOfRanges valueRanges) {
1230 storage.reserve(valueRanges.size());
1231 ranges.reserve(valueRanges.size());
1232 for (ArrayRef<ValueRange> valueRangeGroup : valueRanges) {
1233 storage.push_back(flattenConvertedValues(valueRangeGroup));
1234 }
1235 for (const SmallVector<Value> &values : storage) {
1236 ranges.push_back(values);
1237 }
1238 }
1239};
1240
1243template <typename ValueRangeLike>
1244inline static bool allValueTypesUnifyWithTypes(const ValueRangeLike &values, ArrayRef<Type> types) {
1245 return llvm::all_of_zip(values, types, [](auto value, Type type) {
1246 return typesUnify(value.getType(), type);
1247 });
1248}
1249
1255static ArrayType getSplitPodArrayStorageType(ArrayType arrTy, ArrayRef<StringAttr> recordChain) {
1256 auto elemPodTy = llvm::cast<PodType>(arrTy.getElementType());
1257 Type leafType = getFlattenedTypeAlongPath(elemPodTy, recordChain);
1258 return flattenArrayElementType(arrTy, replaceAffineMapArrayDimsWithWildcards(leafType));
1259}
1260
1262static Value createSplitPodArrayReplacement(
1263 Operation *src, Location loc, ArrayType originalArrTy, const RecordChain &id,
1264 ArrayType preciseSplitType, ConversionPatternRewriter &rewriter,
1265 ArrayRef<ValueRange> mapOperands = {}, DenseI32ArrayAttr numDimsPerMap = nullptr
1266) {
1267 ArrayType storageSplitType = getSplitPodArrayStorageType(originalArrTy, id.nameList);
1268 CreateArrayOp splitArrayOp =
1269 mapOperands.empty()
1270 ? rewriter.create<CreateArrayOp>(loc, storageSplitType)
1271 : rewriter.create<CreateArrayOp>(loc, storageSplitType, mapOperands, numDimsPerMap);
1272 preserveDiscardableAttrs(src, splitArrayOp);
1273 return castValueToTypeIfNeeded(rewriter, loc, splitArrayOp, preciseSplitType);
1274}
1275
1280inline static Value createWritableArrayValue(
1281 OpBuilder &bldr, Location loc, ArrayType arrTy,
1282 std::optional<ArrayInstantiationInfo> instantiation = std::nullopt
1283) {
1284 if (hasAffineMapAttr(arrTy)) {
1285 SmallVector<AffineMapAttr> topLevelMapAttrs;
1286 for (Attribute dimSize : arrTy.getDimensionSizes()) {
1287 if (auto mapAttr = llvm::dyn_cast<AffineMapAttr>(dimSize)) {
1288 topLevelMapAttrs.push_back(mapAttr);
1289 }
1290 }
1291
1292 if (instantiation && !topLevelMapAttrs.empty() &&
1293 countRecursiveAffineMapAttrs(arrTy) == topLevelMapAttrs.size() &&
1294 instantiation->mapOperandStorage.size() == topLevelMapAttrs.size() &&
1295 instantiation->numDimsPerMap.size() == topLevelMapAttrs.size()) {
1296 SmallVector<ValueRange> mapOperands;
1297 mapOperands.reserve(instantiation->mapOperandStorage.size());
1298 for (const SmallVector<Value> &group : instantiation->mapOperandStorage) {
1299 mapOperands.push_back(group);
1300 }
1301 return bldr.create<CreateArrayOp>(loc, arrTy, mapOperands, instantiation->numDimsPerMap);
1302 }
1303
1304 return bldr.create<NonDetOp>(loc, arrTy);
1305 } else {
1306 return bldr.create<CreateArrayOp>(loc, arrTy);
1307 }
1308}
1309
1311static std::optional<Value>
1312tryMaterializeFreshUnwrittenDirectRecordRead(OpBuilder &bldr, Location loc, ReadPodOp readOp) {
1313 if (!isFreshUnwrittenPodRead(readOp)) {
1314 return std::nullopt;
1315 }
1316
1317 Type recordType = readOp.getType();
1318 if (llvm::isa<PodType>(recordType)) {
1319 return std::nullopt;
1320 }
1321
1322 if (ArrayType arrTy = llvm::dyn_cast<ArrayType>(recordType)) {
1323 return createWritableArrayValue(
1324 bldr, loc, arrTy, tryGetFreshUnwrittenPodReadInstantiationInfo(readOp)
1325 );
1326 }
1327
1328 return bldr.create<NonDetOp>(loc, recordType).getResult();
1329}
1330
1332static bool equivalentArrayInstantiationInfo(
1333 const ArrayInstantiationInfo &lhs, const ArrayInstantiationInfo &rhs
1334) {
1335 if (lhs.numDimsPerMap != rhs.numDimsPerMap ||
1336 lhs.mapOperandStorage.size() != rhs.mapOperandStorage.size()) {
1337 return false;
1338 }
1339
1340 for (auto [lhsGroup, rhsGroup] : llvm::zip_equal(lhs.mapOperandStorage, rhs.mapOperandStorage)) {
1341 if (lhsGroup != rhsGroup) {
1342 return false;
1343 }
1344 }
1345
1346 return true;
1347}
1348
1350enum class CommonArrayInstantiationStatus : std::uint8_t {
1351 unavailable,
1352 inferred,
1353 conflict,
1354};
1355
1361static CommonArrayInstantiationStatus
1362inferCommonArrayInstantiation(ArrayRef<Value> values, ArrayInstantiationInfo &result) {
1363 bool initialized = false;
1364 for (Value value : values) {
1365 std::optional<ArrayInstantiationInfo> info = tryGetArrayInstantiationInfo(value);
1366 if (!info) {
1367 return CommonArrayInstantiationStatus::unavailable;
1368 }
1369
1370 if (!initialized) {
1371 result = std::move(*info);
1372 initialized = true;
1373 continue;
1374 }
1375
1376 if (!equivalentArrayInstantiationInfo(result, *info)) {
1377 return CommonArrayInstantiationStatus::conflict;
1378 }
1379 }
1380
1381 return initialized ? CommonArrayInstantiationStatus::inferred
1382 : CommonArrayInstantiationStatus::unavailable;
1383}
1384
1386static Operation *
1387genArrayWrite(OpBuilder &bldr, Location loc, Value arrayRef, ValueRange indices, Value value) {
1388 ArrayType arrTy = llvm::cast<ArrayType>(arrayRef.getType());
1389 Type selectedType = arrTy.getSelectionType(indices.size());
1390 Value convertedValue = castValueToTypeIfNeeded(bldr, loc, value, selectedType);
1391 if (llvm::isa<ArrayType>(selectedType)) {
1392 return bldr.create<InsertArrayOp>(loc, arrayRef, indices, convertedValue);
1393 }
1394 return bldr.create<WriteArrayOp>(loc, arrayRef, indices, convertedValue);
1395}
1396
1397inline static Operation *
1398genArrayWrite(OpBuilder &bldr, Location loc, Value arrayRef, ArrayAttr index, Value value) {
1399 SmallVector<Value> indices = ArrayAccessOpInterface::genIndexConstants(bldr, loc, index);
1400 return genArrayWrite(bldr, loc, arrayRef, indices, value);
1401}
1402
1404static bool tryCollectDirectConvertedPodArrayValues(
1405 Value arrayValue, ArrayType arrTy, ArrayRef<Type> convertedTypes,
1406 SmallVectorImpl<Value> &convertedValues
1407) {
1408 arrayValue = peelUnifiableCasts(arrayValue);
1409
1410 if (auto cast = arrayValue.getDefiningOp<UnrealizedConversionCastOp>()) {
1411 if (cast->getNumResults() != 1 || cast.getResult(0).getType() != arrTy ||
1412 cast->getNumOperands() != convertedTypes.size()) {
1413 return false;
1414 }
1415
1417 cast.getOperands(), convertedTypes, convertedValues
1418 );
1419 }
1420
1421 if (ReadPodOp readOp = arrayValue.getDefiningOp<ReadPodOp>()) {
1422 if (WritePodOp writeOp = findNearestForwardableWrite(readOp)) {
1423 return tryCollectDirectConvertedPodArrayValues(
1424 writeOp.getValue(), arrTy, convertedTypes, convertedValues
1425 );
1426 }
1427 }
1428
1429 return false;
1430}
1431
1433static Value tryCollectDirectConvertedPodArrayShapeSource(
1434 Value arrayValue, ArrayType arrTy, Location loc, OpBuilder &bldr
1435) {
1436 SmallVector<RecordChain> splitIds;
1437 SmallVector<Type> convertedTypes;
1438 collectConvertedPodArrayRecordInfos(arrTy, splitIds, convertedTypes);
1439
1440 SmallVector<Value> convertedValues;
1441 if (!tryCollectDirectConvertedPodArrayValues(
1442 arrayValue, arrTy, convertedTypes, convertedValues
1443 )) {
1444 return {};
1445 }
1446
1447 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
1448 return castValueToTypeIfNeeded(bldr, loc, carrier, getPodArrayShapeCarrierType(arrTy));
1449 }
1450 if (Value leaf = getRankPreservingConvertedPodArrayLeaf(arrTy, convertedValues)) {
1451 return leaf;
1452 }
1453
1454 return {};
1455}
1456
1458static bool hasEarlierWriteToPodInBlock(Operation *op, Value podRef) {
1459 for (Operation &candidate : *op->getBlock()) {
1460 if (&candidate == op) {
1461 return false;
1462 }
1463 if (auto writeOp = dyn_cast<WritePodOp>(&candidate)) {
1464 if (writeOp.getPodRef() == podRef) {
1465 return true;
1466 }
1467 } else if (hasNestedWriteToPod(candidate, podRef)) {
1468 return true;
1469 }
1470 }
1471 return false;
1472}
1473
1475static bool hasEarlierWriteToPod(Operation *op, Value podRef) {
1476 for (Operation *cursor = op; cursor && cursor->getBlock();) {
1477 if (hasEarlierWriteToPodInBlock(cursor, podRef)) {
1478 return true;
1479 }
1480 cursor = cursor->getBlock()->getParentOp();
1481 }
1482 return false;
1483}
1484
1486static bool isFreshUnwrittenPodArrayRead(Value value) {
1487 value = peelUnifiableCasts(value);
1488 ReadPodOp readOp = value.getDefiningOp<ReadPodOp>();
1489 return readOp && splittablePodArray(readOp.getType()) && isFreshUnwrittenPodRead(readOp);
1490}
1491
1493static Value genReadAlongPath(
1494 OpBuilder &bldr, Location loc, Value value, ArrayRef<StringAttr> recordChain,
1495 bool syntheticShapeCarrier
1496) {
1497 if (recordChain.empty()) {
1498 if (ReadPodOp readOp = peelUnifiableCasts(value).getDefiningOp<ReadPodOp>()) {
1499 if (std::optional<Value> materialized =
1500 tryMaterializeFreshUnwrittenDirectRecordRead(bldr, loc, readOp)) {
1501 return *materialized;
1502 }
1503 }
1504 return value;
1505 }
1506
1507 Type valueType = value.getType();
1508 if (llvm::isa<PodType>(valueType)) {
1509 Value nextValue = genRead(bldr, loc, value, recordChain.front());
1510 return genReadAlongPath(bldr, loc, nextValue, recordChain.drop_front(), syntheticShapeCarrier);
1511 }
1512
1513 if (ArrayType arrTy = splittablePodArray(valueType)) {
1514 Type splitType = getFlattenedTypeAlongPath(valueType, recordChain, syntheticShapeCarrier);
1515 auto splitArrTy = llvm::dyn_cast<ArrayType>(splitType);
1516 assert(splitArrTy);
1517
1518 SmallVector<RecordChain> splitIds;
1519 SmallVector<Type> splitTypes;
1520 collectConvertedPodArrayRecordInfos(arrTy, splitIds, splitTypes);
1521 auto *splitIt = llvm::find(splitIds, RecordChain(recordChain, syntheticShapeCarrier));
1522 assert(splitIt != splitIds.end() && "record path must name a flattened POD array leaf");
1523 size_t splitIdx = std::distance(splitIds.begin(), splitIt);
1524
1525 SmallVector<Value> convertedValues;
1526 if (tryCollectDirectConvertedPodArrayValues(value, arrTy, splitTypes, convertedValues)) {
1527 return convertedValues[splitIdx];
1528 }
1529
1530 Value strippedValue = peelUnifiableCasts(value);
1531 if (isFreshUnwrittenPodArrayRead(value)) {
1532 return createWritableArrayValue(bldr, loc, splitArrTy, tryGetArrayInstantiationInfo(value));
1533 }
1534
1535 if (strippedValue.getDefiningOp<ReadPodOp>()) {
1536 auto splitReads =
1537 bldr.create<UnrealizedConversionCastOp>(loc, TypeRange(splitTypes), strippedValue);
1538 return splitReads.getResult(splitIdx);
1539 }
1540
1541 if (syntheticShapeCarrier) {
1542 assert(
1543 isPodArrayShapeCarrierMarker(recordChain.back()) &&
1544 "synthetic shape carrier must use the reserved shape marker"
1545 );
1546 return bldr.create<CreateArrayOp>(loc, splitArrTy);
1547 }
1548
1549 if (!arrTy.hasStaticShape()) {
1550 llvm_unreachable(
1551 "non-static nested array-of-POD scalarization requires split-array backing or an "
1552 "uninitialized pod field"
1553 );
1554 }
1555
1556 auto subIndices = arrTy.getSubelementIndices();
1557 assert(subIndices && "static-shape arrays must provide subelement indices");
1558
1559 Value splitArray = createWritableArrayValue(bldr, loc, splitArrTy);
1560 for (ArrayAttr index : *subIndices) {
1561 Value element = ArrayAccessOpInterface::genRead(bldr, loc, value, index);
1562 Value leafValue = genReadAlongPath(bldr, loc, element, recordChain, syntheticShapeCarrier);
1563 genArrayWrite(bldr, loc, splitArray, index, leafValue);
1564 }
1565 return splitArray;
1566 }
1567
1568 llvm_unreachable("record path cannot continue through a non-POD leaf");
1569}
1570
1572inline static Value
1573genReadAlongPath(OpBuilder &bldr, Location loc, Value podRef, const RecordChain &recordChain) {
1574 return genReadAlongPath(
1575 bldr, loc, podRef, ArrayRef(recordChain.nameList), recordChain.syntheticShapeCarrier
1576 );
1577}
1578
1580static SmallVector<Value> materializeCompatiblePodArrayLeafValues(
1581 Location loc, Value arrayValue, ArrayType arrTy, OpBuilder &rewriter
1582) {
1583 SmallVector<RecordChain> splitIds;
1584 SmallVector<Type> splitTypes;
1585 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
1586
1587 SmallVector<Value> leaves;
1588 leaves.reserve(splitIds.size());
1589 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
1590 Value splitValue = genReadAlongPath(rewriter, loc, arrayValue, id);
1591 leaves.push_back(castValueToTypeIfNeeded(rewriter, loc, splitValue, splitType));
1592 }
1593 return leaves;
1594}
1595
1596using VirtualPodLeafMap = DenseMap<RecordChain, Value>;
1597using VirtualPodValueMap = DenseMap<Value, VirtualPodLeafMap>;
1598
1600static Value rebuildFlattenedPodRecord(
1601 OpBuilder &bldr, Location loc, Type recordType, SmallVectorImpl<StringAttr> &recordChain,
1602 const VirtualPodLeafMap &leafValues
1603) {
1604 if (PodType nestedPodTy = dyn_cast<PodType>(recordType)) {
1605 NewPodOp nestedPod = bldr.create<NewPodOp>(loc, nestedPodTy);
1606 for (RecordAttr record : nestedPodTy.getRecords()) {
1607 recordChain.push_back(record.getName());
1608 Value recordValue =
1609 rebuildFlattenedPodRecord(bldr, loc, record.getType(), recordChain, leafValues);
1610 genWrite(bldr, loc, nestedPod, record.getName(), recordValue);
1611 recordChain.pop_back();
1612 }
1613 return nestedPod;
1614 }
1615
1616 if (ArrayType arrTy = splittablePodArray(recordType)) {
1617 if (!arrTy.hasStaticShape()) {
1618 SmallVector<RecordChain> splitIds;
1619 SmallVector<Type> splitTypes;
1620 collectConvertedPodArrayRecordInfos(arrTy, splitIds, splitTypes);
1621
1622 SmallVector<Value> leafArrays;
1623 leafArrays.reserve(splitIds.size());
1624 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
1625 auto it = leafValues.find(id.withPrefix(recordChain));
1626 assert(it != leafValues.end() && "missing flattened POD array leaf value");
1627 leafArrays.push_back(castValueToTypeIfNeeded(bldr, loc, it->second, splitType));
1628 }
1629
1630 return bldr.create<UnrealizedConversionCastOp>(loc, TypeRange {arrTy}, leafArrays)
1631 .getResult(0);
1632 }
1633
1634 auto elemPodTy = llvm::cast<PodType>(arrTy.getElementType());
1635 auto subIndices = arrTy.getSubelementIndices();
1636 assert(subIndices && "static-shape arrays must provide subelement indices");
1637
1638 Value rebuiltArray = bldr.create<CreateArrayOp>(loc, arrTy);
1639 for (ArrayAttr index : *subIndices) {
1640 VirtualPodLeafMap elementLeafValues;
1641 SmallVector<StringAttr> elementRecordChain;
1642 forEachPodLeaf(elemPodTy, elementRecordChain, [&](const RecordChain &id, Type) {
1643 auto it = leafValues.find(id.withPrefix(recordChain));
1644 assert(it != leafValues.end() && "missing flattened POD array leaf value");
1645 elementLeafValues[id] = ArrayAccessOpInterface::genRead(bldr, loc, it->second, index);
1646 });
1647
1648 NewPodOp elementPod = bldr.create<NewPodOp>(loc, elemPodTy);
1649 SmallVector<StringAttr> nestedChain;
1650 for (RecordAttr record : elemPodTy.getRecords()) {
1651 nestedChain.push_back(record.getName());
1652 Value recordValue =
1653 rebuildFlattenedPodRecord(bldr, loc, record.getType(), nestedChain, elementLeafValues);
1654 genWrite(bldr, loc, elementPod, record.getName(), recordValue);
1655 nestedChain.pop_back();
1656 }
1657 genArrayWrite(bldr, loc, rebuiltArray, index, elementPod);
1658 }
1659 return rebuiltArray;
1660 }
1661
1662 auto it = leafValues.find(RecordChain(recordChain));
1663 assert(it != leafValues.end() && "missing flattened POD leaf value");
1664 return it->second;
1665}
1666
1668static Value peelVirtualPodCompatibilityCasts(Value value) {
1669 while (auto cast = value.getDefiningOp<UnifiableCastOp>()) {
1670 value = cast.getInput();
1671 }
1672 return value;
1673}
1674
1676static const VirtualPodLeafMap *
1677lookupVirtualPodLeafMap(Value podValue, const VirtualPodValueMap &virtualPods) {
1678 podValue = peelVirtualPodCompatibilityCasts(podValue);
1679 auto it = virtualPods.find(podValue);
1680 return it != virtualPods.end() ? &it->second : nullptr;
1681}
1682
1684inline static VirtualPodValueMap::iterator
1685lookupVirtualPodLeafMapIt(Value podValue, VirtualPodValueMap &virtualPods) {
1686 return virtualPods.find(peelVirtualPodCompatibilityCasts(podValue));
1687}
1688
1690static SmallVector<Value> orderedVirtualPodLeafValues(
1691 PodType podTy, Location loc, OpBuilder &bldr, const VirtualPodLeafMap &leafValues
1692) {
1693 SmallVector<Value> orderedValues;
1694 SmallVector<StringAttr> recordChain;
1695 forEachPodLeaf(
1696 podTy, recordChain,
1697 [&leafValues, &orderedValues, &bldr, loc](const RecordChain &id, Type leafType) {
1698 auto it = leafValues.find(id);
1699 assert(it != leafValues.end() && "missing virtual POD leaf value");
1700 orderedValues.push_back(castValueToTypeIfNeeded(bldr, loc, it->second, leafType));
1701 }
1702 );
1703 return orderedValues;
1704}
1705
1706// If the operand has PodType, add reads from all pod records to the `newOperands` list otherwise
1707// add the original operand to the list.
1708static void processInputOperand(
1709 Location loc, Value operand, SmallVector<Value> &newOperands, OpBuilder &rewriter,
1710 Operation *userOp = nullptr, const VirtualPodValueMap *virtualPods = nullptr
1711) {
1712 if (PodType pt = splittablePod(operand.getType())) {
1713 if (virtualPods) {
1714 if (const VirtualPodLeafMap *leafValues = lookupVirtualPodLeafMap(operand, *virtualPods);
1715 leafValues && (!userOp || !hasEarlierWriteToPod(userOp, operand))) {
1716 llvm::append_range(
1717 newOperands, orderedVirtualPodLeafValues(pt, loc, rewriter, *leafValues)
1718 );
1719 return;
1720 }
1721 }
1722 SmallVector<StringAttr> recordChain;
1723 forEachPodLeaf(pt, recordChain, [&](const RecordChain &id, Type) {
1724 newOperands.push_back(genReadAlongPath(rewriter, loc, operand, id));
1725 });
1726 } else {
1727 newOperands.push_back(operand);
1728 }
1729}
1730
1733static void processInputOperands(
1734 ValueRange operands, MutableOperandRange outputOpRef, Operation *op,
1735 ConversionPatternRewriter &rewriter, const VirtualPodValueMap *virtualPods = nullptr
1736) {
1737 SmallVector<Value> newOperands;
1738 for (Value v : operands) {
1739 processInputOperand(op->getLoc(), v, newOperands, rewriter, op, virtualPods);
1740 }
1741 rewriter.modifyOpInPlace(op, [&outputOpRef, &newOperands]() {
1742 outputOpRef.assign(ValueRange(newOperands));
1743 });
1744}
1745
1747static PodType getWholePodEqualityType(constrain::EmitEqualityOp op) {
1748 if (PodType lhsTy = splittablePod(peelUnifiableCasts(op.getLhs()).getType())) {
1749 return lhsTy;
1750 }
1751 return splittablePod(peelUnifiableCasts(op.getRhs()).getType());
1752}
1753
1755static void setInsertionPointAfterValueDefinition(Value value, OpBuilder &bldr) {
1756 if (Operation *defOp = value.getDefiningOp()) {
1757 bldr.setInsertionPointAfter(defOp);
1758 } else {
1759 auto blockArg = llvm::cast<BlockArgument>(value);
1760 bldr.setInsertionPointToStart(blockArg.getOwner());
1761 }
1762}
1763
1765template <typename EmitValuesFn>
1766inline static void materializeCompatibleValuesAfterDefinition(
1767 Location loc, Value source, ArrayRef<Type> targetTypes, OpBuilder &rewriter,
1768 SmallVectorImpl<Value> &out, EmitValuesFn &&emitValues
1769) {
1770 if (!targetTypes.empty()) {
1771 OpBuilder::InsertionGuard guard(rewriter);
1772 setInsertionPointAfterValueDefinition(source, rewriter);
1773 emitValues(loc, source, targetTypes, rewriter, out);
1774 }
1775}
1776
1778static Value
1779getCompatibleMaterializationSource(Value source, Type compatibleType, const char *assertMessage) {
1780 source = peelUnifiableCasts(source);
1781 assert(typesUnify(source.getType(), compatibleType) && assertMessage);
1782 return source;
1783}
1784
1787template <typename CollectTypesFn>
1788static ArrayRef<Value> getOrMaterializeCompatibleLeafValues(
1789 Location location, Value source, Type shapeTy, OpBuilder &rewriter,
1790 CompatiblePodLeafMaterializationMap &materializedLeaves, CollectTypesFn &&collectTypes,
1791 const char *assertMessage
1792) {
1793 source = getCompatibleMaterializationSource(source, shapeTy, assertMessage);
1794 CompatiblePodLeafMaterializationKey key {source, shapeTy};
1795 auto [it, inserted] = materializedLeaves.try_emplace(key);
1796 if (!inserted) {
1797 return it->second;
1798 }
1799
1800 SmallVector<Type> leafTypes;
1801 collectTypes(leafTypes);
1802 materializeCompatibleValuesAfterDefinition(
1803 location, source, leafTypes, rewriter, it->second,
1804 [](Location loc, Value src, ArrayRef<Type> targetTypes, OpBuilder &bldr,
1805 SmallVectorImpl<Value> &out) {
1806 auto splitCast = bldr.create<UnrealizedConversionCastOp>(loc, TypeRange(targetTypes), src);
1807 llvm::append_range(out, splitCast.getResults());
1808 }
1809 );
1810 return it->second;
1811}
1812
1819static ArrayRef<Value> getOrMaterializeCompatiblePodLeafValues(
1820 Location loc, Value source, PodType podTy, OpBuilder &rewriter,
1821 CompatiblePodLeafMaterializationMap &materializedLeaves
1822) {
1823 return getOrMaterializeCompatibleLeafValues(
1824 loc, source, podTy, rewriter, materializedLeaves, [&podTy](SmallVector<Type> &leafTypes) {
1825 splitPodTypeTo(podTy, leafTypes);
1826 }, "materialized POD leaves require a source type compatible with the POD shape"
1827 );
1828}
1829
1831static ArrayRef<Value> getOrMaterializeCompatiblePodArrayLeafValues(
1832 Location loc, Value source, ArrayType arrTy, OpBuilder &rewriter,
1833 CompatiblePodLeafMaterializationMap &materializedLeaves
1834) {
1835 return getOrMaterializeCompatibleLeafValues(
1836 loc, source, arrTy, rewriter, materializedLeaves, [&arrTy](SmallVector<Type> &leafTypes) {
1837 splitPodArrayTypeTo(arrTy, leafTypes);
1838 }, "materialized POD-array leaves require a source type compatible with the array shape"
1839 );
1840}
1841
1844static SmallVector<Value> materializeCompatiblePodArrayConvertedValues(
1845 Location location, Value source, ArrayType arrTy, OpBuilder &rewriter
1846) {
1847 SmallVector<RecordChain> splitIds;
1848 SmallVector<Type> convertedTypes;
1849 collectConvertedPodArrayRecordInfos(arrTy, splitIds, convertedTypes);
1850 source = getCompatibleMaterializationSource(
1851 source, arrTy,
1852 "materialized POD-array components require a source type compatible with the "
1853 "array shape"
1854 );
1855
1856 SmallVector<Value> convertedValues;
1857 materializeCompatibleValuesAfterDefinition(
1858 location, source, convertedTypes, rewriter, convertedValues,
1859 [](Location loc, Value src, ArrayRef<Type> targetTypes, OpBuilder &bldr,
1860 SmallVectorImpl<Value> &out) {
1861 if (targetTypes.size() == 1) {
1862 out.push_back(castValueToTypeIfNeeded(bldr, loc, src, targetTypes.front()));
1863 return;
1864 }
1865 // Keep the compatible components tied to one split cast so function arguments can
1866 // be promoted to the concrete component types instead of casting one argument
1867 // independently to each type.
1868 auto splitCast = bldr.create<UnrealizedConversionCastOp>(loc, TypeRange(targetTypes), src);
1869 llvm::append_range(out, splitCast.getResults());
1870 }
1871 );
1872 return convertedValues;
1873}
1874
1876static void collectWholePodEqualityOperandLeaves(
1877 Location loc, Value operand, PodType podTy, SmallVector<Value> &leaves, OpBuilder &rewriter,
1878 CompatiblePodLeafMaterializationMap &materializedLeaves, Operation *userOp = nullptr,
1879 const VirtualPodValueMap *virtualPods = nullptr
1880) {
1881 operand = peelUnifiableCasts(operand);
1882 if (splittablePod(operand.getType())) {
1883 processInputOperand(loc, operand, leaves, rewriter, userOp, virtualPods);
1884 return;
1885 }
1886
1887 assert(
1888 typesUnify(operand.getType(), podTy) &&
1889 "whole-POD equality operand must unify with the concrete POD shape"
1890 );
1891 llvm::append_range(
1892 leaves,
1893 getOrMaterializeCompatiblePodLeafValues(loc, operand, podTy, rewriter, materializedLeaves)
1894 );
1895}
1896
1898static LogicalResult splitWholePodEmitEquality(
1899 constrain::EmitEqualityOp op, RewriterBase &rewriter,
1900 CompatiblePodLeafMaterializationMap &materializedLeaves, Operation *userOp = nullptr,
1901 const VirtualPodValueMap *virtualPods = nullptr
1902) {
1903 PodType podTy = getWholePodEqualityType(op);
1904 assert(podTy && "expected a whole-POD equality to expose a concrete POD side");
1905
1906 SmallVector<Value> lhsLeaves;
1907 SmallVector<Value> rhsLeaves;
1908 collectWholePodEqualityOperandLeaves(
1909 op.getLoc(), op.getLhs(), podTy, lhsLeaves, rewriter, materializedLeaves, userOp, virtualPods
1910 );
1911 collectWholePodEqualityOperandLeaves(
1912 op.getLoc(), op.getRhs(), podTy, rhsLeaves, rewriter, materializedLeaves, userOp, virtualPods
1913 );
1914
1915 assert(lhsLeaves.size() == rhsLeaves.size() && "POD equality leaves must stay aligned");
1916 for (auto [lhsLeaf, rhsLeaf] : llvm::zip_equal(lhsLeaves, rhsLeaves)) {
1917 if (failed(rejectRaggedNestedLeafEquality(op, lhsLeaf, rhsLeaf))) {
1918 return failure();
1919 }
1920 preserveDiscardableAttrs(
1921 op, rewriter.create<constrain::EmitEqualityOp>(op.getLoc(), lhsLeaf, rhsLeaf)
1922 );
1923 }
1924 rewriter.eraseOp(op);
1925 return success();
1926}
1927
1934static Value createVirtualPodPlaceholder(
1935 OpBuilder &bldr, Location loc, PodType podTy, const VirtualPodLeafMap &leafValues
1936) {
1937 if (!hasAffineMapAttr(podTy)) {
1938 return bldr.create<NewPodOp>(loc, podTy);
1939 }
1940
1941 SmallVector<Value> orderedValues = orderedVirtualPodLeafValues(podTy, loc, bldr, leafValues);
1942 return bldr.create<UnrealizedConversionCastOp>(loc, TypeRange {podTy}, orderedValues)
1943 .getResult(0);
1944}
1945
1947inline static void
1948materializeVirtualPod(OpBuilder &bldr, NewPodOp pod, const VirtualPodLeafMap &leafValues) {
1949 Location loc = pod.getLoc();
1950 PodType podTy = pod.getType();
1951 SmallVector<StringAttr> recordChain;
1952 for (RecordAttr record : podTy.getRecords()) {
1953 recordChain.push_back(record.getName());
1954 Value recordValue =
1955 rebuildFlattenedPodRecord(bldr, loc, record.getType(), recordChain, leafValues);
1956 genWrite(bldr, loc, pod, record.getName(), recordValue);
1957 recordChain.pop_back();
1958 }
1959}
1960
1967static Operation *
1968findVirtualPodMaterializationAnchor(NewPodOp pod, const VirtualPodLeafMap &leafValues) {
1969 Operation *anchor = pod.getOperation();
1970 Block *block = anchor->getBlock();
1971
1972 for (const auto &it : leafValues) {
1973 Operation *defOp = it.second.getDefiningOp();
1974 if (!defOp || defOp->getBlock() != block) {
1975 continue;
1976 }
1977 if (anchor->isBeforeInBlock(defOp)) {
1978 anchor = defOp;
1979 }
1980 }
1981
1982 return anchor;
1983}
1984
1989inline static bool isInsideSupportedScfRegion(Operation *op) {
1991}
1992
1994static bool canResolveVirtualPodRead(ReadPodOp op, const VirtualPodValueMap &virtualPods) {
1995 if (!lookupVirtualPodLeafMap(op.getPodRef(), virtualPods) || hasEarlierWrite(op) ||
1996 findNearestForwardableWrite(op)) {
1997 return false;
1998 }
1999 Type recType = llvm::cast<PodType>(op.getPodRefType()).getRecordMap().lookup(op.getRecordName());
2000 return llvm::isa<PodType>(recType) || !splittablePodArray(recType);
2001}
2002
2004inline static bool shouldDeferPodArrayReadToStep3(ReadArrayOp op) {
2005 return splittablePodArray(op.getArrRefType()) &&
2006 llvm::isa_and_present<ReadPodOp>(op.getArrRef().getDefiningOp());
2007}
2008
2010static ReadPodOp getReadPodBacking(Value value) {
2011 value = peelUnifiableCasts(value);
2012 if (ReadPodOp readOp = value.getDefiningOp<ReadPodOp>()) {
2013 return readOp;
2014 }
2015
2016 auto cast = value.getDefiningOp<UnrealizedConversionCastOp>();
2017 if (!cast || cast->getNumOperands() != 1) {
2018 return {};
2019 }
2020 return peelUnifiableCasts(cast.getOperand(0)).getDefiningOp<ReadPodOp>();
2021}
2022
2024inline static bool shouldDeferPodArrayLengthToStep3(ArrayLengthOp op) {
2025 return splittablePodArray(op.getArrRefType()) && getReadPodBacking(op.getArrRef());
2026}
2027
2029static SmallVector<std::string> getSplitRecordNameSuffixes(Type type) {
2030 SmallVector<std::string> suffixes;
2031 if (PodType pt = splittablePod(type)) {
2032 SmallVector<StringAttr> recordChain;
2033 forEachPodLeaf(pt, recordChain, [&suffixes](const RecordChain &id, Type) {
2034 std::string suffix;
2035 llvm::raw_string_ostream os(suffix);
2036 for (size_t i = 0; i < id.nameList.size(); ++i) {
2037 os << '.' << id.getComponentName(i);
2038 }
2039 suffixes.push_back(std::move(suffix));
2040 });
2041 }
2042 return suffixes;
2043}
2044
2046static void updateVirtualPodRecordLeafValues(
2047 Location loc, StringAttr recordName, Type recordType, Value recordValue,
2048 const VirtualPodValueMap &virtualPods, RewriterBase &rewriter, VirtualPodLeafMap &leafValues
2049) {
2050 SmallVector<StringAttr> prefix {recordName};
2051
2052 if (PodType nestedPodTy = llvm::dyn_cast<PodType>(recordType)) {
2053 if (const VirtualPodLeafMap *nestedLeafValues =
2054 lookupVirtualPodLeafMap(recordValue, virtualPods)) {
2055 SmallVector<StringAttr> nestedRecordChain;
2056 forEachPodLeaf(nestedPodTy, nestedRecordChain, [&](const RecordChain &id, Type) {
2057 leafValues[id.withPrefix(prefix)] = nestedLeafValues->at(id);
2058 });
2059 return;
2060 }
2061
2062 SmallVector<StringAttr> nestedRecordChain;
2063 forEachPodLeaf(nestedPodTy, nestedRecordChain, [&](const RecordChain &id, Type) {
2064 leafValues[id.withPrefix(prefix)] = genReadAlongPath(rewriter, loc, recordValue, id);
2065 });
2066 return;
2067 }
2068
2069 if (ArrayType arrTy = splittablePodArray(recordType)) {
2070 SmallVector<RecordChain> splitIds;
2071 SmallVector<Type> splitTypes;
2072 collectConvertedPodArrayRecordInfos(arrTy, splitIds, splitTypes);
2073 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2074 leafValues[id.withPrefix(prefix)] = castValueToTypeIfNeeded(
2075 rewriter, loc, genReadAlongPath(rewriter, loc, recordValue, id), splitType
2076 );
2077 }
2078 return;
2079 }
2080
2081 leafValues[RecordChain(prefix)] = castValueToTypeIfNeeded(rewriter, loc, recordValue, recordType);
2082}
2083
2085inline static void baseTargetSetup(ConversionTarget &target) {
2086 target.addLegalDialect<
2091 scf::SCFDialect>();
2092 target.addLegalOp<ModuleOp>();
2093}
2094
2097class NondetToNewPod : public OpConversionPattern<NonDetOp> {
2098 using OpConversionPattern<NonDetOp>::OpConversionPattern;
2099 LogicalResult matchAndRewrite(
2100 NonDetOp nondetOp, OpAdaptor, ConversionPatternRewriter &rewriter
2101 ) const override {
2102 if (auto pt = dyn_cast<PodType>(nondetOp.getType())) {
2103 preserveDiscardableAttrs(nondetOp, rewriter.replaceOpWithNewOp<NewPodOp>(nondetOp, pt));
2104 return success();
2105 }
2106 return failure();
2107 }
2108};
2109
2112static LogicalResult step0(ModuleOp modOp) {
2113 MLIRContext *ctx = modOp.getContext();
2114
2115 PassManager prepPM(ctx);
2118 if (failed(prepPM.run(modOp))) {
2119 return failure();
2120 }
2121
2122 RewritePatternSet patterns {ctx};
2123 patterns.add<NondetToNewPod>(ctx);
2124 ConversionTarget target {*ctx};
2125
2126 baseTargetSetup(target);
2127 target.addLegalOp<UnrealizedConversionCastOp>();
2128 target.addDynamicallyLegalOp<NonDetOp>([](NonDetOp op) { return !isa<PodType>(op.getType()); });
2129
2130 return applyFullConversion(modOp, target, std::move(patterns));
2131}
2132
2134using MemberInfo = std::pair<StringAttr, Type>;
2136using LocalMemberReplacementMap = DenseMap<RecordChain, MemberInfo>;
2138using MemberReplacementMap = DenseMap<StructDefOp, DenseMap<StringAttr, LocalMemberReplacementMap>>;
2139
2141inline static StringAttr getSplitPodArrayShapeMemberName(MLIRContext *ctx, StringAttr memberName) {
2142 return StringAttr::get(ctx, (memberName.getValue() + "_shape").str());
2143}
2144
2146static void flattenPodMemberIntoLeaves(
2147 MemberDefOp originalMember, PodType podTy, SmallVectorImpl<StringAttr> &recordChain,
2148 LocalMemberReplacementMap &localRepMapRef, SymbolTable &structSymbolTable,
2149 ConversionPatternRewriter &rewriter
2150) {
2151 forEachPodLeaf(podTy, recordChain, [&](const RecordChain &id, Type ty) {
2152 StringAttr name =
2153 id.getFlattenedMemberName(originalMember.getContext(), originalMember.getSymNameAttr());
2154 MemberDefOp newMember = rewriter.create<MemberDefOp>(
2155 originalMember.getLoc(), name, ty, !id.syntheticShapeCarrier && originalMember.getSignal(),
2156 !id.syntheticShapeCarrier && originalMember.getColumn()
2157 );
2158 preserveDiscardableAttrs(originalMember, newMember);
2159 newMember.setPublicAttr(originalMember.hasPublicAttr());
2160 localRepMapRef[id] = std::make_pair(structSymbolTable.insert(newMember), ty);
2161 });
2162}
2163
2168class SplitPodInMemberDefOp : public OpConversionPattern<MemberDefOp> {
2169 SymbolTableCollection &tables;
2170 MemberReplacementMap &repMapRef;
2171
2172public:
2173 SplitPodInMemberDefOp(
2174 MLIRContext *ctx, SymbolTableCollection &symTables, MemberReplacementMap &memberRepMap
2175 )
2176 : OpConversionPattern<MemberDefOp>(ctx), tables(symTables), repMapRef(memberRepMap) {}
2177
2178 inline static bool legal(MemberDefOp op) { return !splittablePod(op.getType()); }
2179
2180 LogicalResult match(MemberDefOp op) const override { return failure(legal(op)); }
2181
2182 void
2183 rewrite(MemberDefOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
2184 StructDefOp inStruct = op->getParentOfType<StructDefOp>();
2185 assert(inStruct);
2186 LocalMemberReplacementMap &localRepMapRef = repMapRef[inStruct][op.getSymNameAttr()];
2187
2188 PodType podTy = llvm::cast<PodType>(adaptor.getType()); // safe per legal() check
2189
2190 SymbolTable &structSymbolTable = tables.getSymbolTable(inStruct);
2191 SmallVector<StringAttr> recordChain;
2192 flattenPodMemberIntoLeaves(op, podTy, recordChain, localRepMapRef, structSymbolTable, rewriter);
2193 rewriter.eraseOp(op);
2194 }
2195};
2196
2198class SplitPodArrayInMemberDefOp : public OpConversionPattern<MemberDefOp> {
2199 SymbolTableCollection &tables;
2200 MemberReplacementMap &repMapRef;
2201
2202public:
2203 SplitPodArrayInMemberDefOp(
2204 MLIRContext *ctx, SymbolTableCollection &symTables, MemberReplacementMap &memberRepMap
2205 )
2206 : OpConversionPattern<MemberDefOp>(ctx), tables(symTables), repMapRef(memberRepMap) {}
2207
2208 inline static bool legal(MemberDefOp op) { return !splittablePodArray(op.getType()); }
2209
2210 LogicalResult match(MemberDefOp op) const override { return failure(legal(op)); }
2211
2212 void
2213 rewrite(MemberDefOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
2214 StructDefOp inStruct = op->getParentOfType<StructDefOp>();
2215 assert(inStruct);
2216 LocalMemberReplacementMap &localRepMapRef = repMapRef[inStruct][op.getSymNameAttr()];
2217
2218 ArrayType arrTy = llvm::cast<ArrayType>(adaptor.getType());
2219 SmallVector<RecordChain> splitIds;
2220 SmallVector<Type> splitTypes;
2221 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2222 if (splitTypes.empty()) {
2223 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2224 rewriter.modifyOpInPlace(op, [&]() {
2225 op.setType(carrierTy);
2226 op.removeSignalAttr();
2227 });
2228 localRepMapRef[RecordChain()] = std::make_pair(op.getSymNameAttr(), carrierTy);
2229 return;
2230 }
2231
2232 SymbolTable &structSymbolTable = tables.getSymbolTable(inStruct);
2233 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2234 StringAttr name = id.getFlattenedMemberName(op.getContext(), op.getSymNameAttr());
2235 MemberDefOp newMember = rewriter.create<MemberDefOp>(
2236 op.getLoc(), name, splitType, op.getSignal(), op.getColumn()
2237 );
2238 preserveDiscardableAttrs(op, newMember);
2239 newMember.setPublicAttr(op.hasPublicAttr());
2240 localRepMapRef[id] = std::make_pair(structSymbolTable.insert(newMember), splitType);
2241 }
2242 if (needsPodArrayShapeCarrier(arrTy)) {
2243 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2244 StringAttr carrierName =
2245 getSplitPodArrayShapeMemberName(op.getContext(), op.getSymNameAttr());
2246 MemberDefOp carrierMember =
2247 rewriter.create<MemberDefOp>(op.getLoc(), carrierName, carrierTy, false, op.getColumn());
2248 preserveDiscardableAttrs(op, carrierMember);
2249 carrierMember.setPublicAttr(op.hasPublicAttr());
2250 localRepMapRef[RecordChain()] =
2251 std::make_pair(structSymbolTable.insert(carrierMember), carrierTy);
2252 }
2253 rewriter.eraseOp(op);
2254 }
2255};
2256
2259static LogicalResult
2260step1(ModuleOp modOp, SymbolTableCollection &symTables, MemberReplacementMap &memberRepMap) {
2261 MLIRContext *ctx = modOp.getContext();
2262
2263 RewritePatternSet patterns(ctx);
2264
2265 patterns.add<SplitPodInMemberDefOp, SplitPodArrayInMemberDefOp>(ctx, symTables, memberRepMap);
2266
2267 ConversionTarget target(*ctx);
2268 baseTargetSetup(target);
2269 target.addLegalOp<UnrealizedConversionCastOp>();
2270 target.addDynamicallyLegalOp<MemberDefOp>([](MemberDefOp op) {
2271 return SplitPodInMemberDefOp::legal(op) && SplitPodArrayInMemberDefOp::legal(op);
2272 });
2273
2274 LLVM_DEBUG(llvm::dbgs() << "Begin step 1: split pod-type and array-of-pod members\n";);
2275 return applyFullConversion(modOp, target, std::move(patterns));
2276}
2277
2283class PodArrayTypeConverter : public TypeConverter {
2284public:
2285 PodArrayTypeConverter() {
2286 addConversion([](Type type) { return type; });
2287 addConversion(
2288 [](ArrayType arrTy, SmallVectorImpl<Type> &results) -> std::optional<LogicalResult> {
2289 if (!splittablePodArray(arrTy)) {
2290 return std::nullopt;
2291 }
2292 convertPodArrayTypeTo(arrTy, results);
2293 return success();
2294 }
2295 );
2296
2297 auto materializeCast = [](OpBuilder &bldr, Type targetType, ValueRange inputs,
2298 Location loc) -> Value {
2299 if (inputs.size() != 1 || !typesUnify(inputs.front().getType(), targetType)) {
2300 return {};
2301 }
2302 return castValueToTypeIfNeeded(bldr, loc, inputs.front(), targetType);
2303 };
2304 addTargetMaterialization(materializeCast);
2305 addArgumentMaterialization(materializeCast);
2306 addSourceMaterialization(materializeCast);
2307 }
2308};
2309
2311class SplitPodArrayNonDetOp : public OpConversionPattern<NonDetOp> {
2312public:
2313 using OpConversionPattern<NonDetOp>::OpConversionPattern;
2314
2315 static bool legal(NonDetOp op) { return !splittablePodArray(op.getType()); }
2316
2317 LogicalResult match(NonDetOp op) const override { return failure(legal(op)); }
2318
2319 void rewrite(NonDetOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
2320 SmallVector<Type> splitTypes;
2321 splitPodArrayTypeTo(op.getType(), splitTypes);
2322 if (splitTypes.empty()) {
2323 preserveDiscardableAttrs(
2324 op, rewriter.replaceOpWithNewOp<NonDetOp>(
2325 op, getPodArrayShapeCarrierType(llvm::cast<ArrayType>(op.getType()))
2326 )
2327 );
2328 return;
2329 }
2330 SmallVector<Value> replacements;
2331 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
2332 replacements.reserve(splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
2333 for (Type splitType : splitTypes) {
2334 replacements.push_back(
2335 preserveDiscardableAttrs(op, rewriter.create<NonDetOp>(op.getLoc(), splitType))
2336 );
2337 }
2338 if (needsPodArrayShapeCarrier(arrTy)) {
2339 replacements.push_back(preserveDiscardableAttrs(
2340 op, rewriter.create<NonDetOp>(op.getLoc(), getPodArrayShapeCarrierType(arrTy))
2341 ));
2342 }
2343 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2344 }
2345};
2346
2357class SplitPodArrayCreateArrayOp : public OpConversionPattern<CreateArrayOp> {
2358public:
2359 using OpConversionPattern<CreateArrayOp>::OpConversionPattern;
2360
2361 static bool legal(CreateArrayOp op) { return !splittablePodArray(op.getType()); }
2362
2363 LogicalResult matchAndRewrite(
2364 CreateArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2365 ) const override {
2366 if (legal(op)) {
2367 return failure();
2368 }
2369 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
2370 SmallVector<RecordChain> splitIds;
2371 SmallVector<Type> splitTypes;
2372 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2373 if (splitTypes.empty()) {
2374 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2375 if (adaptor.getMapOperands().empty()) {
2376 preserveDiscardableAttrs(op, rewriter.replaceOpWithNewOp<CreateArrayOp>(op, carrierTy));
2377 return success();
2378 }
2379
2380 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2381 preserveDiscardableAttrs(
2382 op, rewriter.replaceOpWithNewOp<CreateArrayOp>(
2383 op, carrierTy, mapOperands.ranges, op.getNumDimsPerMapAttr()
2384 )
2385 );
2386 return success();
2387 }
2388
2389 SmallVector<Value> replacements;
2390 replacements.reserve(splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
2391 DenseI32ArrayAttr numDimsPerMap = op.getNumDimsPerMapAttr();
2392 if (isNullOrEmpty(numDimsPerMap)) {
2393 if (adaptor.getElements().empty()) {
2394 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2395 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2396 replacements.push_back(
2397 createSplitPodArrayReplacement(op, op.getLoc(), arrTy, id, preciseSplitType, rewriter)
2398 );
2399 }
2400 if (needsPodArrayShapeCarrier(arrTy)) {
2401 Value shapeCarrier =
2402 materializeArrayLengthCarrier(op.getResult(), arrTy, op.getLoc(), rewriter);
2403 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
2404 replacements.push_back(shapeCarrier);
2405 }
2406 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2407 return success();
2408 }
2409
2410 auto elementIndices = arrTy.getSubelementIndices();
2411 assert(elementIndices && "array.new with explicit elements requires a static array shape");
2412 assert(
2413 elementIndices->size() == adaptor.getElements().size() &&
2414 "array.new element count must match the outer array cardinality"
2415 );
2416
2417 // Inline initializers are linearized only across the original outer array dimensions. When
2418 // a flattened POD leaf is itself an array, populate the rewritten split array one outer
2419 // element at a time so each leaf array becomes a subarray insert rather than a malformed
2420 // inline operand to the flattened `array.new`.
2421 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2422 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2423 ArrayType storageSplitType = getSplitPodArrayStorageType(arrTy, id.nameList);
2424
2425 SmallVector<Value> leafValues;
2426 leafValues.reserve(adaptor.getElements().size());
2427 for (ValueRange elementRange : adaptor.getElements()) {
2428 Value element = getSingleConvertedValue(elementRange);
2429 leafValues.push_back(genReadAlongPath(rewriter, op.getLoc(), element, id));
2430 }
2431
2432 ArrayType materializedType = storageSplitType;
2433 Value splitArray;
2434 if (storageSplitType != preciseSplitType) {
2435 ArrayInstantiationInfo instantiationInfo;
2436 switch (inferCommonArrayInstantiation(leafValues, instantiationInfo)) {
2437 case CommonArrayInstantiationStatus::conflict:
2438 // TODO: this POD could be promoted to a complete `struct.def` but that's not easy.
2439 op.emitOpError(
2440 "with POD elements having conflicting affine map instantiations cannot be promoted "
2441 "to higher dimensional array"
2442 );
2443 return failure();
2444 case CommonArrayInstantiationStatus::inferred: {
2445 materializedType = preciseSplitType;
2446 SmallVector<ValueRange> mapOperands;
2447 mapOperands.reserve(instantiationInfo.mapOperandStorage.size());
2448 for (const SmallVector<Value> &values : instantiationInfo.mapOperandStorage) {
2449 mapOperands.push_back(values);
2450 }
2451 CreateArrayOp splitArrayOp = rewriter.create<CreateArrayOp>(
2452 op.getLoc(), materializedType, mapOperands, instantiationInfo.numDimsPerMap
2453 );
2454 preserveDiscardableAttrs(op, splitArrayOp);
2455 splitArray = splitArrayOp;
2456 break;
2457 }
2458 case CommonArrayInstantiationStatus::unavailable:
2459 break;
2460 }
2461 }
2462
2463 if (!splitArray) {
2464 splitArray = createWritableArrayValue(rewriter, op.getLoc(), materializedType);
2465 preserveDiscardableAttrs(op, splitArray.getDefiningOp());
2466 }
2467
2468 for (auto [index, leafValue] : llvm::zip_equal(*elementIndices, leafValues)) {
2469 genArrayWrite(rewriter, op.getLoc(), splitArray, index, leafValue);
2470 }
2471 replacements.push_back(
2472 castValueToTypeIfNeeded(rewriter, op.getLoc(), splitArray, preciseSplitType)
2473 );
2474 }
2475 } else {
2476 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2477 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2478 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2479 replacements.push_back(createSplitPodArrayReplacement(
2480 op, op.getLoc(), arrTy, id, preciseSplitType, rewriter, mapOperands.ranges,
2481 numDimsPerMap
2482 ));
2483 }
2484 }
2485
2486 if (needsPodArrayShapeCarrier(arrTy)) {
2487 Value shapeCarrier =
2488 materializeArrayLengthCarrier(op.getResult(), arrTy, op.getLoc(), rewriter);
2489 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
2490 replacements.push_back(shapeCarrier);
2491 }
2492
2493 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2494 return success();
2495 }
2496};
2497
2499class SplitPodArrayReadArrayOp : public OpConversionPattern<ReadArrayOp> {
2500public:
2501 using OpConversionPattern<ReadArrayOp>::OpConversionPattern;
2502
2503 static bool legal(ReadArrayOp op) {
2504 return !splittablePodArray(op.getArrRefType()) || shouldDeferPodArrayReadToStep3(op);
2505 }
2506
2507 LogicalResult matchAndRewrite(
2508 ReadArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2509 ) const override {
2510 if (legal(op)) {
2511 return failure();
2512 }
2513 ArrayType arrTy = op.getArrRefType();
2514 PodType podTy = llvm::cast<PodType>(arrTy.getElementType());
2515 SmallVector<RecordChain> splitIds;
2516 SmallVector<Type> splitTypes;
2517 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2518 if (splitTypes.empty()) {
2519 preserveDiscardableAttrs(op, rewriter.replaceOpWithNewOp<NewPodOp>(op, podTy));
2520 return success();
2521 }
2522
2523 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
2524 NewPodOp pod = rewriter.create<NewPodOp>(op.getLoc(), podTy);
2525 preserveDiscardableAttrs(op, pod);
2526 VirtualPodLeafMap leafValues;
2527 auto splitArrRefs = adaptor.getArrRef().take_front(splitIds.size());
2528 for (auto [id, splitType, splitArrRange] :
2529 llvm::zip_equal(splitIds, splitTypes, splitArrRefs)) {
2530 Value leafValue = ArrayAccessOpInterface::genRead(
2531 rewriter, op.getLoc(), getSingleConvertedValue(splitArrRange), indices
2532 );
2533 preserveDiscardableAttrs(op, leafValue.getDefiningOp());
2534 leafValues[id] = tagRaggedNestedLeafValue(
2535 rewriter, op.getLoc(), leafValue, getRaggedNestedLeafAttrName(arrTy, splitType)
2536 );
2537 }
2538
2539 SmallVector<StringAttr> recordChain;
2540 for (RecordAttr record : podTy.getRecords()) {
2541 recordChain.push_back(record.getName());
2542 Value recordValue = rebuildFlattenedPodRecord(
2543 rewriter, op.getLoc(), record.getType(), recordChain, leafValues
2544 );
2545 genWrite(rewriter, op.getLoc(), pod, record.getName(), recordValue);
2546 recordChain.pop_back();
2547 }
2548 rewriter.replaceOp(op, pod);
2549 return success();
2550 }
2551};
2552
2554class SplitPodArrayWriteArrayOp : public OpConversionPattern<WriteArrayOp> {
2555public:
2556 using OpConversionPattern<WriteArrayOp>::OpConversionPattern;
2557
2558 static bool legal(WriteArrayOp op) { return !splittablePodArray(op.getArrRefType()); }
2559
2560 LogicalResult matchAndRewrite(
2561 WriteArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2562 ) const override {
2563 if (legal(op)) {
2564 return failure();
2565 }
2566 ArrayType arrTy = op.getArrRefType();
2567 SmallVector<RecordChain> splitIds;
2568 SmallVector<Type> splitTypes;
2569 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2570 if (splitTypes.empty()) {
2571 rewriter.eraseOp(op);
2572 return success();
2573 }
2574
2575 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
2576 Value podValue = getSingleConvertedValue(adaptor.getRvalue());
2577 auto splitArrRefs = adaptor.getArrRef().take_front(splitIds.size());
2578 for (auto [id, splitArrRange, splitType] :
2579 llvm::zip_equal(splitIds, splitArrRefs, splitTypes)) {
2580 Value leafValue = genReadAlongPath(rewriter, op.getLoc(), podValue, id);
2581 preserveDiscardableAttrs(
2582 op, genArrayWrite(
2583 rewriter, op.getLoc(), getSingleConvertedValue(splitArrRange), indices, leafValue
2584 )
2585 );
2586 }
2587 rewriter.eraseOp(op);
2588 return success();
2589 }
2590};
2591
2593class SplitPodArrayInFuncDefOp : public OpConversionPattern<FuncDefOp> {
2594public:
2595 using OpConversionPattern<FuncDefOp>::OpConversionPattern;
2596
2597 static bool legal(FuncDefOp op) {
2598 return !containsSplittablePodArrayType(op.getArgumentTypes()) &&
2599 !containsSplittablePodArrayType(op.getResultTypes());
2600 }
2601
2602 LogicalResult match(FuncDefOp op) const override { return failure(legal(op)); }
2603
2604 LogicalResult
2605 matchAndRewrite(FuncDefOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
2606 const auto *tyConv = getTypeConverter();
2607 assert(tyConv && "expected pod-array type converter");
2608
2609 FunctionType oldTy = op.getFunctionType();
2610 TypeConverter::SignatureConversion inputConversion(oldTy.getNumInputs());
2611 if (failed(tyConv->convertSignatureArgs(oldTy.getInputs(), inputConversion))) {
2612 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod inputs");
2613 }
2614
2615 SmallVector<Type> newResults;
2616 if (failed(tyConv->convertTypes(oldTy.getResults(), newResults))) {
2617 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod results");
2618 }
2619
2620 if (!op.getBody().empty() &&
2621 failed(rewriter.convertRegionTypes(&op.getBody(), *tyConv, &inputConversion))) {
2622 return rewriter.notifyMatchFailure(op, "failed to convert function body block arguments");
2623 }
2624
2625 SmallVector<size_t> originalInputIdxToSize, originalResultIdxToSize;
2626 SmallVector<Type> newInputs = convertPodArrayTypes(oldTy.getInputs(), &originalInputIdxToSize);
2627 SmallVector<Type> newResultsWithSizeInfo =
2628 convertPodArrayTypes(oldTy.getResults(), &originalResultIdxToSize);
2629 assert(
2630 newResultsWithSizeInfo == newResults &&
2631 "expected array-of-pod type conversion to match function result attr replication"
2632 );
2633 SplitFunctionNameInfo inputNameInfo =
2634 collectSplitFunctionNameInfo(op.getArgumentTypes(), [&](unsigned i) {
2635 return op.getArgNameAttr(i);
2636 }, getSplitPodArrayRecordNameSuffixes);
2637 ArrayAttr resultAttrs = op.getAllResultAttrs();
2638 SplitFunctionNameInfo resultNameInfo =
2639 collectSplitFunctionNameInfo(op.getResultTypes(), [resultAttrs](unsigned i) {
2640 return getAttrAtIndexWithName(resultAttrs, i, RES_NAME_ATTR_NAME);
2641 }, getSplitPodArrayRecordNameSuffixes);
2642
2643 rewriter.modifyOpInPlace(op, [&]() {
2644 op.setFunctionType(FunctionType::get(op.getContext(), newInputs, newResults));
2645 if (ArrayAttr newArgAttrs = replicateFunctionNameAttrsAsNeeded(
2646 op.getArgAttrsAttr(), originalInputIdxToSize, newInputs, ARG_NAME_ATTR_NAME,
2647 inputNameInfo.originalNames, inputNameInfo.existingNames,
2648 inputNameInfo.splitNameSuffixes
2649 )) {
2650 op.setArgAttrsAttr(newArgAttrs);
2651 }
2652 if (ArrayAttr newResAttrs = replicateFunctionNameAttrsAsNeeded(
2653 op.getResAttrsAttr(), originalResultIdxToSize, newResults, RES_NAME_ATTR_NAME,
2654 resultNameInfo.originalNames, resultNameInfo.existingNames,
2655 resultNameInfo.splitNameSuffixes
2656 )) {
2657 op.setResAttrsAttr(newResAttrs);
2658 }
2659 });
2660 return success();
2661 }
2662};
2663
2670static void collectSplitPodArrayOperandValues(
2671 Location loc, Value originalOperand, ValueRange convertedValues,
2672 SmallVectorImpl<Value> &newOperands, ConversionPatternRewriter &rewriter
2673) {
2674 ArrayType arrTy = splittablePodArray(originalOperand.getType());
2675 if (!arrTy) {
2676 llvm::append_range(newOperands, convertedValues);
2677 return;
2678 }
2679
2680 SmallVector<RecordChain> splitIds;
2681 SmallVector<Type> splitTypes;
2682 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2683 if (splitTypes.empty()) {
2684 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
2685 newOperands.push_back(
2686 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2687 );
2688 return;
2689 }
2690 if (!convertedValues.empty()) {
2691 newOperands.push_back(castValueToTypeIfNeeded(
2692 rewriter, loc, getSingleConvertedValue(convertedValues),
2693 getPodArrayShapeCarrierType(arrTy)
2694 ));
2695 return;
2696 }
2697 newOperands.push_back(materializeArrayLengthCarrier(originalOperand, arrTy, loc, rewriter));
2698 return;
2699 }
2700
2701 ValueRange leafConvertedValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2702 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues);
2703
2704 auto isDirectAggregateToSplitCast = [&leafConvertedValues, &splitTypes]() {
2705 if (leafConvertedValues.empty()) {
2706 return false;
2707 }
2708 auto castOp = leafConvertedValues.front().getDefiningOp<UnrealizedConversionCastOp>();
2709 if (!castOp || castOp->getNumOperands() != 1) {
2710 return false;
2711 }
2712 ArrayType castArrTy = splittablePodArray(castOp.getOperand(0).getType());
2713 size_t expectedResults =
2714 castArrTy ? splitTypes.size() + (needsPodArrayShapeCarrier(castArrTy) ? 1 : 0) : 0;
2715 if (!castArrTy || castOp->getNumResults() != expectedResults) {
2716 return false;
2717 }
2718
2719 return llvm::all_of(llvm::zip_equal(leafConvertedValues, splitTypes), [&castOp](auto pair) {
2720 Value convertedValue = std::get<0>(pair);
2721 Type splitType = std::get<1>(pair);
2722 return convertedValue.getDefiningOp<UnrealizedConversionCastOp>() == castOp &&
2723 typesUnify(convertedValue.getType(), splitType);
2724 });
2725 };
2726 bool directAggregateToSplitCast = isDirectAggregateToSplitCast();
2727
2728 if (!directAggregateToSplitCast && allValueTypesUnifyWithTypes(leafConvertedValues, splitTypes)) {
2729 llvm::append_range(newOperands, leafConvertedValues);
2730 if (carrier) {
2731 newOperands.push_back(
2732 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2733 );
2734 } else if (needsPodArrayShapeCarrier(arrTy)) {
2735 newOperands.push_back(materializeArrayLengthCarrier(originalOperand, arrTy, loc, rewriter));
2736 }
2737 return;
2738 }
2739
2740 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2741 Value splitValue = genReadAlongPath(rewriter, loc, originalOperand, id);
2742 newOperands.push_back(castValueToTypeIfNeeded(rewriter, loc, splitValue, splitType));
2743 }
2744 if (carrier && !directAggregateToSplitCast) {
2745 newOperands.push_back(
2746 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2747 );
2748 } else if (needsPodArrayShapeCarrier(arrTy)) {
2749 RecordChain carrierId({getPodArrayShapeCarrierMarker(rewriter.getContext())}, true);
2750 Value splitCarrier = genReadAlongPath(rewriter, loc, originalOperand, carrierId);
2751 newOperands.push_back(
2752 castValueToTypeIfNeeded(rewriter, loc, splitCarrier, getPodArrayShapeCarrierType(arrTy))
2753 );
2754 }
2755}
2756
2758class SplitPodArrayInUnifiableCastOp : public OpConversionPattern<UnifiableCastOp> {
2759public:
2760 using OpConversionPattern<UnifiableCastOp>::OpConversionPattern;
2761
2762 static bool legal(UnifiableCastOp op) {
2763 return !splittablePodArray(op.getType()) && !splittablePodArray(op.getInput().getType());
2764 }
2765
2766 LogicalResult matchAndRewrite(
2767 UnifiableCastOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2768 ) const override {
2769 if (legal(op)) {
2770 return failure();
2771 }
2772
2773 ArrayType inputArrTy = splittablePodArray(op.getInput().getType());
2774 ArrayType resultArrTy = splittablePodArray(op.getType());
2775
2776 // Lowering a split array-of-POD input to one non-array SSA value would require either
2777 // materializing the original aggregate or rewriting the surrounding function/template
2778 // signature to thread separate generic pieces.
2779 if (inputArrTy && !resultArrTy) {
2780 return rewriter.notifyMatchFailure(
2781 op, "array-of-pod input to non-array result requires aggregate materialization or "
2782 "signature/template rewriting"
2783 );
2784 }
2785
2786 if (!inputArrTy) {
2787 return rewriter.notifyMatchFailure(
2788 op, "generic input to array-of-pod result requires signature/template rewriting"
2789 );
2790 }
2791
2792 SmallVector<RecordChain> inputSplitIds;
2793 SmallVector<Type> inputSplitTypes;
2794 splitPodArrayTypeTo(inputArrTy, inputSplitTypes, &inputSplitIds);
2795
2796 SmallVector<RecordChain> resultSplitIds;
2797 SmallVector<Type> resultSplitTypes;
2798 splitPodArrayTypeTo(resultArrTy, resultSplitTypes, &resultSplitIds);
2799
2800 if (inputSplitIds != resultSplitIds) {
2801 return rewriter.notifyMatchFailure(
2802 op, "array-of-pod cast changed POD leaf structure unexpectedly"
2803 );
2804 }
2805
2806 SmallVector<Value> splitInputs;
2807 collectSplitPodArrayOperandValues(
2808 op.getLoc(), op.getInput(), adaptor.getInput(), splitInputs, rewriter
2809 );
2810 ValueRange splitInputLeaves = getConvertedPodArrayLeafValues(inputArrTy, splitInputs);
2811 if (resultSplitTypes.empty()) {
2812 if (splitInputs.size() != 1) {
2813 return rewriter.notifyMatchFailure(
2814 op, "expected one shape carrier for zero-leaf array-of-pod cast"
2815 );
2816 }
2817 Value replacement = castValueToTypeIfNeeded(
2818 rewriter, op.getLoc(), splitInputs.front(), getPodArrayShapeCarrierType(resultArrTy)
2819 );
2820 if (replacement != splitInputs.front()) {
2821 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2822 }
2823 rewriter.replaceOp(op, replacement);
2824 return success();
2825 }
2826 if (splitInputLeaves.size() != resultSplitTypes.size()) {
2827 return rewriter.notifyMatchFailure(
2828 op, "failed to collect one split input per array-of-pod cast leaf"
2829 );
2830 }
2831
2832 SmallVector<Value> replacements;
2833 replacements.reserve(
2834 resultSplitTypes.size() + (needsPodArrayShapeCarrier(resultArrTy) ? 1 : 0)
2835 );
2836 for (auto [splitInput, resultSplitType] : llvm::zip_equal(splitInputLeaves, resultSplitTypes)) {
2837 Value replacement =
2838 castValueToTypeIfNeeded(rewriter, op.getLoc(), splitInput, resultSplitType);
2839 if (replacement != splitInput) {
2840 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2841 }
2842 replacements.push_back(replacement);
2843 }
2844 if (needsPodArrayShapeCarrier(resultArrTy)) {
2845 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(inputArrTy, splitInputs);
2846 if (!carrier) {
2847 carrier = materializeArrayLengthCarrier(op.getInput(), inputArrTy, op.getLoc(), rewriter);
2848 }
2849 Value replacement = castValueToTypeIfNeeded(
2850 rewriter, op.getLoc(), carrier, getPodArrayShapeCarrierType(resultArrTy)
2851 );
2852 if (replacement != carrier) {
2853 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2854 }
2855 replacements.push_back(replacement);
2856 }
2857
2858 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2859 return success();
2860 }
2861};
2862
2864class SplitPodArrayInReturnOp : public OpConversionPattern<ReturnOp> {
2865public:
2866 using OpConversionPattern<ReturnOp>::OpConversionPattern;
2867
2868 static bool legal(ReturnOp op) {
2869 return !containsSplittablePodArrayType(op.getOperands().getTypes());
2870 }
2871
2872 LogicalResult matchAndRewrite(
2873 ReturnOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2874 ) const override {
2875 if (legal(op)) {
2876 return failure();
2877 }
2878 SmallVector<Value> newOperands;
2879 for (auto [operand, convertedValues] :
2880 llvm::zip_equal(op.getOperands(), adaptor.getOperands())) {
2881 collectSplitPodArrayOperandValues(
2882 op.getLoc(), operand, convertedValues, newOperands, rewriter
2883 );
2884 }
2885 preserveDiscardableAttrs(
2886 op, rewriter.replaceOpWithNewOp<ReturnOp>(op, ValueRange(newOperands))
2887 );
2888 return success();
2889 }
2890};
2891
2893class SplitPodArrayInCallOp : public OpConversionPattern<CallOp> {
2894public:
2895 using OpConversionPattern<CallOp>::OpConversionPattern;
2896
2897 static bool legal(CallOp op) {
2898 return !containsSplittablePodArrayType(op.getArgOperands().getTypes()) &&
2899 !containsSplittablePodArrayType(op.getResultTypes());
2900 }
2901
2902 LogicalResult match(CallOp op) const override { return failure(legal(op)); }
2903
2904 LogicalResult matchAndRewrite(
2905 CallOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2906 ) const override {
2907 const auto *tyConv = getTypeConverter();
2908 assert(tyConv && "expected pod-array type converter");
2909
2910 SmallVector<Type> newResultTypes;
2911 if (failed(tyConv->convertTypes(op.getResultTypes(), newResultTypes))) {
2912 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod call results");
2913 }
2914
2915 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2916
2917 SmallVector<Value> newArgOperands;
2918 for (auto [operand, convertedValues] :
2919 llvm::zip_equal(op.getArgOperands(), adaptor.getArgOperands())) {
2920 collectSplitPodArrayOperandValues(
2921 op.getLoc(), operand, convertedValues, newArgOperands, rewriter
2922 );
2923 }
2925 op.getLoc(), newResultTypes, op, mapOperands.ranges, newArgOperands, rewriter
2926 );
2927
2928 SmallVector<SmallVector<Value>> replacementStorage;
2929 replacementStorage.reserve(op.getNumResults());
2930 auto newResultIt = newCall.getResults().begin();
2931 for (Type oldResultType : op.getResultTypes()) {
2932 SmallVector<Type> convertedTypes;
2933 (void)convertPodArrayTypeTo(oldResultType, convertedTypes);
2934 SmallVector<Value> replacementsForResult;
2935 replacementsForResult.reserve(convertedTypes.size());
2936 for (size_t i = 0; i < convertedTypes.size(); ++i) {
2937 replacementsForResult.push_back(*newResultIt);
2938 ++newResultIt;
2939 }
2940 replacementStorage.push_back(std::move(replacementsForResult));
2941 }
2942
2943 SmallVector<ValueRange> replacements;
2944 replacements.reserve(replacementStorage.size());
2945 for (const SmallVector<Value> &values : replacementStorage) {
2946 replacements.push_back(values);
2947 }
2948 rewriter.replaceOpWithMultiple(op, replacements);
2949 return success();
2950 }
2951};
2952
2954static Value getPodArrayEqualityShapeSource(
2955 ArrayType arrTy, Value originalValue, ValueRange convertedValues, Location loc,
2956 OpBuilder &rewriter
2957) {
2958 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
2959 return carrier;
2960 }
2961
2962 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2963 if (!leafValues.empty()) {
2964 return leafValues.front();
2965 }
2966
2967 if (!convertedValues.empty()) {
2968 return getSingleConvertedValue(convertedValues);
2969 }
2970
2971 return materializeArrayLengthCarrier(originalValue, arrTy, loc, rewriter);
2972}
2973
2975static SmallVector<Value> collectCompatiblePodArrayEqualityLeaves(
2976 Location loc, Value originalValue, ValueRange convertedValues, ArrayType arrTy,
2977 ArrayType peerArrTy, ConversionPatternRewriter &rewriter,
2978 CompatiblePodLeafMaterializationMap &materializedLeaves
2979) {
2980 if (arrTy) {
2981 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2982 size_t leafCount = getSplitPodArrayLeafCount(arrTy);
2983 if (leafValues.size() == leafCount || leafCount == 0) {
2984 return SmallVector<Value>(leafValues.begin(), leafValues.end());
2985 }
2986
2987 return materializeCompatiblePodArrayLeafValues(loc, originalValue, arrTy, rewriter);
2988 }
2989
2990 if (!peerArrTy || convertedValues.size() != 1) {
2991 if (!peerArrTy || !convertedValues.empty()) {
2992 return SmallVector<Value>(convertedValues.begin(), convertedValues.end());
2993 }
2994 }
2995
2996 Value source = convertedValues.empty() ? originalValue : getSingleConvertedValue(convertedValues);
2997 SmallVector<Value> materialized;
2998 llvm::append_range(
2999 materialized, getOrMaterializeCompatiblePodArrayLeafValues(
3000 loc, source, peerArrTy, rewriter, materializedLeaves
3001 )
3002 );
3003 return materialized;
3004}
3005
3007static ArrayType getCompatiblePodArrayType(ArrayType arrTy, PodType elemPodTy) {
3008 if (ArrayType concreteArrTy = splittablePodArray(arrTy)) {
3009 return concreteArrTy;
3010 }
3011 return ArrayType::get(elemPodTy, arrTy.getDimensionSizes());
3012}
3013
3015class SplitPodArrayInEmitEqualityOp : public OpConversionPattern<constrain::EmitEqualityOp> {
3016 CompatiblePodLeafMaterializationMap &materializedLeaves;
3017
3018public:
3019 SplitPodArrayInEmitEqualityOp(
3020 TypeConverter &tyConv, MLIRContext *ctx,
3021 CompatiblePodLeafMaterializationMap &materializedLeafMap
3022 )
3023 : OpConversionPattern<constrain::EmitEqualityOp>(tyConv, ctx),
3024 materializedLeaves(materializedLeafMap) {}
3025
3026 static bool legal(constrain::EmitEqualityOp op) {
3027 return !containsSplittablePodArrayType(op->getOperandTypes());
3028 }
3029
3030 LogicalResult matchAndRewrite(
3031 constrain::EmitEqualityOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3032 ) const override {
3033 if (legal(op)) {
3034 return failure();
3035 }
3036
3037 ArrayType lhsTy = splittablePodArray(op.getLhs().getType());
3038 ArrayType rhsTy = splittablePodArray(op.getRhs().getType());
3039 StringRef raggedKind = lhsTy ? getRaggedNestedLeafKind(lhsTy) : StringRef {};
3040 if (raggedKind.empty() && rhsTy) {
3041 raggedKind = getRaggedNestedLeafKind(rhsTy);
3042 }
3043 if (!raggedKind.empty()) {
3044 return op.emitOpError()
3045 << "cannot lower nested " << raggedKind
3046 << " array leaf equality for array-of-POD without per-element shape witnesses";
3047 }
3048
3049 bool lhsNeedsShapeCheck = lhsTy && needsPodArrayShapeCarrier(lhsTy);
3050 bool rhsNeedsShapeCheck = rhsTy && needsPodArrayShapeCarrier(rhsTy);
3051 SmallVector<Value> lhsCompatibleConvertedValues;
3052 SmallVector<Value> rhsCompatibleConvertedValues;
3053 if (rhsTy && rhsNeedsShapeCheck && !lhsTy &&
3054 (adaptor.getLhs().empty() || adaptor.getLhs().size() == 1)) {
3055 Value lhsSource =
3056 adaptor.getLhs().empty() ? op.getLhs() : getSingleConvertedValue(adaptor.getLhs());
3057 lhsCompatibleConvertedValues =
3058 materializeCompatiblePodArrayConvertedValues(op.getLoc(), lhsSource, rhsTy, rewriter);
3059 }
3060 if (lhsTy && lhsNeedsShapeCheck && !rhsTy &&
3061 (adaptor.getRhs().empty() || adaptor.getRhs().size() == 1)) {
3062 Value rhsSource =
3063 adaptor.getRhs().empty() ? op.getRhs() : getSingleConvertedValue(adaptor.getRhs());
3064 rhsCompatibleConvertedValues =
3065 materializeCompatiblePodArrayConvertedValues(op.getLoc(), rhsSource, lhsTy, rewriter);
3066 }
3067
3068 ArrayType shapeCheckTy = lhsTy ? lhsTy : rhsTy;
3069 bool needsShapeCheck = lhsTy ? lhsNeedsShapeCheck : rhsNeedsShapeCheck;
3070 if (lhsTy && rhsTy) {
3071 if (lhsTy.getDimensionSizes().size() != rhsTy.getDimensionSizes().size()) {
3072 return rewriter.notifyMatchFailure(
3073 op, "expected array-of-pod equality operands with matching rank"
3074 );
3075 }
3076 needsShapeCheck = lhsNeedsShapeCheck || rhsNeedsShapeCheck;
3077 }
3078
3079 if (shapeCheckTy && needsShapeCheck) {
3080 Value lhsShapeSource = getPodArrayEqualityShapeSource(
3081 lhsTy ? lhsTy : rhsTy, op.getLhs(),
3082 lhsTy ? adaptor.getLhs() : ValueRange(lhsCompatibleConvertedValues), op.getLoc(), rewriter
3083 );
3084 Value rhsShapeSource = getPodArrayEqualityShapeSource(
3085 rhsTy ? rhsTy : lhsTy, op.getRhs(),
3086 rhsTy ? adaptor.getRhs() : ValueRange(rhsCompatibleConvertedValues), op.getLoc(), rewriter
3087 );
3088
3089 for (size_t dim = 0, rank = shapeCheckTy.getDimensionSizes().size(); dim < rank; ++dim) {
3090 Value dimVal = rewriter.create<arith::ConstantOp>(
3091 op.getLoc(), rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3092 );
3093 Value lhsLen = rewriter.create<ArrayLengthOp>(op.getLoc(), lhsShapeSource, dimVal);
3094 Value rhsLen = rewriter.create<ArrayLengthOp>(op.getLoc(), rhsShapeSource, dimVal);
3095 preserveDiscardableAttrs(
3096 op, rewriter.create<constrain::EmitEqualityOp>(op.getLoc(), lhsLen, rhsLen)
3097 );
3098 }
3099 }
3100
3101 SmallVector<Value> lhsLeaves;
3102 if (!lhsTy && rhsTy && !lhsCompatibleConvertedValues.empty()) {
3103 ValueRange materializedLeavesForLhs =
3104 getConvertedPodArrayLeafValues(rhsTy, lhsCompatibleConvertedValues);
3105 llvm::append_range(lhsLeaves, materializedLeavesForLhs);
3106 } else {
3107 lhsLeaves = collectCompatiblePodArrayEqualityLeaves(
3108 op.getLoc(), op.getLhs(), adaptor.getLhs(), lhsTy, rhsTy, rewriter, materializedLeaves
3109 );
3110 }
3111
3112 SmallVector<Value> rhsLeaves;
3113 if (!rhsTy && lhsTy && !rhsCompatibleConvertedValues.empty()) {
3114 ValueRange materializedLeavesForRhs =
3115 getConvertedPodArrayLeafValues(lhsTy, rhsCompatibleConvertedValues);
3116 llvm::append_range(rhsLeaves, materializedLeavesForRhs);
3117 } else {
3118 rhsLeaves = collectCompatiblePodArrayEqualityLeaves(
3119 op.getLoc(), op.getRhs(), adaptor.getRhs(), rhsTy, lhsTy, rewriter, materializedLeaves
3120 );
3121 }
3122 if (lhsLeaves.size() != rhsLeaves.size()) {
3123 return rewriter.notifyMatchFailure(
3124 op, "expected array-of-pod equality operands to expand to the same number of leaves"
3125 );
3126 }
3127
3128 for (auto [lhs, rhs] : llvm::zip_equal(lhsLeaves, rhsLeaves)) {
3129 preserveDiscardableAttrs(
3130 op, rewriter.create<constrain::EmitEqualityOp>(op.getLoc(), lhs, rhs)
3131 );
3132 }
3133 rewriter.eraseOp(op);
3134 return success();
3135 }
3136};
3137
3160class SplitPodArrayInEmitContainmentOp : public OpConversionPattern<constrain::EmitContainmentOp> {
3161 CompatiblePodLeafMaterializationMap &materializedLeaves;
3162
3163public:
3164 SplitPodArrayInEmitContainmentOp(
3165 TypeConverter &tyConv, MLIRContext *ctx,
3166 CompatiblePodLeafMaterializationMap &materializedLeafMap
3167 )
3168 : OpConversionPattern<constrain::EmitContainmentOp>(tyConv, ctx),
3169 materializedLeaves(materializedLeafMap) {}
3170
3171 static bool legal(constrain::EmitContainmentOp op) {
3172 return !containsSplittablePodArrayType(op->getOperandTypes()) &&
3173 getTaggedRaggedNestedLeafKind(op.getLhs()).empty() &&
3174 getTaggedRaggedNestedLeafKind(op.getRhs()).empty();
3175 }
3176
3178 static SmallVector<Value> collectContainmentLeaves(
3179 Location loc, Value originalOperand, ValueRange convertedValues,
3180 ConversionPatternRewriter &rewriter,
3181 CompatiblePodLeafMaterializationMap *materializedLeaves = nullptr,
3182 std::optional<PodType> compatiblePodTy = std::nullopt
3183 ) {
3184 if (ArrayType arrTy = splittablePodArray(originalOperand.getType())) {
3185 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
3186 if (hasZeroLeafPodArraySplit(arrTy)) {
3187 return {};
3188 }
3189 size_t leafCount = getSplitPodArrayLeafCount(arrTy);
3190 if (leafValues.size() == leafCount) {
3191 return SmallVector<Value>(leafValues.begin(), leafValues.end());
3192 }
3193 return materializeCompatiblePodArrayLeafValues(loc, originalOperand, arrTy, rewriter);
3194 }
3195
3196 if (splittablePod(originalOperand.getType())) {
3197 SmallVector<Value> podLeaves;
3198 processInputOperand(loc, getSingleConvertedValue(convertedValues), podLeaves, rewriter);
3199 return podLeaves;
3200 }
3201
3202 if (materializedLeaves && compatiblePodTy &&
3203 (convertedValues.empty() || convertedValues.size() == 1)) {
3204 Value source =
3205 convertedValues.empty() ? originalOperand : getSingleConvertedValue(convertedValues);
3206 ArrayRef<Value> leaves = getOrMaterializeCompatiblePodLeafValues(
3207 loc, source, *compatiblePodTy, rewriter, *materializedLeaves
3208 );
3209 return SmallVector<Value>(leaves.begin(), leaves.end());
3210 }
3211
3212 return SmallVector<Value>(convertedValues.begin(), convertedValues.end());
3213 }
3214
3215 LogicalResult matchAndRewrite(
3216 constrain::EmitContainmentOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3217 ) const override {
3218 if (failed(rejectRaggedNestedLeafContainment(op))) {
3219 return failure();
3220 }
3221
3222 if (legal(op)) {
3223 return failure();
3224 }
3225
3226 Location loc = op.getLoc();
3227 ArrayType lhsTy = op.getLhs().getType();
3228 Type rhsTy = op.getRhs().getType();
3229
3230 size_t lhsRank = lhsTy.getDimensionSizes().size();
3231 size_t rhsRank = 0;
3232 if (auto rhsArrTy = llvm::dyn_cast<ArrayType>(rhsTy)) {
3233 rhsRank = rhsArrTy.getDimensionSizes().size();
3234 }
3235 assert(lhsRank >= rhsRank && "constrain.in verifier should reject higher-rank rhs arrays");
3236 size_t selectedDims = lhsRank - rhsRank;
3237
3238 ArrayType lhsPodArrTy = splittablePodArray(lhsTy);
3239 ArrayType rhsPodArrTy = splittablePodArray(rhsTy);
3240 assert(
3241 (lhsPodArrTy || rhsPodArrTy) &&
3242 "containment rewrite requires at least one concrete array-of-POD operand"
3243 );
3244
3245 PodType compatibleElemPodTy = lhsPodArrTy ? llvm::cast<PodType>(lhsPodArrTy.getElementType())
3246 : llvm::cast<PodType>(rhsPodArrTy.getElementType());
3247 ArrayType compatibleLhsTy = getCompatiblePodArrayType(lhsTy, compatibleElemPodTy);
3248 ArrayType rhsArrTy = llvm::dyn_cast<ArrayType>(rhsTy);
3249 ArrayType compatibleRhsArrTy =
3250 rhsArrTy ? getCompatiblePodArrayType(rhsArrTy, compatibleElemPodTy) : ArrayType();
3251
3252 SmallVector<Value> lhsCompatibleConvertedValues;
3253 if (!lhsPodArrTy) {
3254 if (!adaptor.getLhs().empty() && adaptor.getLhs().size() != 1) {
3255 return rewriter.notifyMatchFailure(
3256 op, "expected a single converted lhs value to materialize generic containment leaves"
3257 );
3258 }
3259 Value lhsSource =
3260 adaptor.getLhs().empty() ? op.getLhs() : getSingleConvertedValue(adaptor.getLhs());
3261 lhsCompatibleConvertedValues =
3262 materializeCompatiblePodArrayConvertedValues(loc, lhsSource, compatibleLhsTy, rewriter);
3263 }
3264
3265 SmallVector<Value> rhsCompatibleConvertedValues;
3266 if (compatibleRhsArrTy && !rhsPodArrTy) {
3267 if (!adaptor.getRhs().empty() && adaptor.getRhs().size() != 1) {
3268 return rewriter.notifyMatchFailure(
3269 op, "expected a single converted rhs value to materialize generic containment leaves"
3270 );
3271 }
3272 Value rhsSource =
3273 adaptor.getRhs().empty() ? op.getRhs() : getSingleConvertedValue(adaptor.getRhs());
3274 rhsCompatibleConvertedValues = materializeCompatiblePodArrayConvertedValues(
3275 loc, rhsSource, compatibleRhsArrTy, rewriter
3276 );
3277 }
3278
3279 SmallVector<Value> lhsLeaves;
3280 if (!lhsCompatibleConvertedValues.empty()) {
3281 llvm::append_range(
3282 lhsLeaves, getConvertedPodArrayLeafValues(compatibleLhsTy, lhsCompatibleConvertedValues)
3283 );
3284 } else {
3285 lhsLeaves = collectContainmentLeaves(loc, op.getLhs(), adaptor.getLhs(), rewriter);
3286 }
3287
3288 SmallVector<Value> rhsLeaves;
3289 if (!rhsCompatibleConvertedValues.empty()) {
3290 llvm::append_range(
3291 rhsLeaves,
3292 getConvertedPodArrayLeafValues(compatibleRhsArrTy, rhsCompatibleConvertedValues)
3293 );
3294 } else {
3295 rhsLeaves = collectContainmentLeaves(
3296 loc, op.getRhs(), adaptor.getRhs(), rewriter, &materializedLeaves, compatibleElemPodTy
3297 );
3298 }
3299 if (lhsLeaves.size() != rhsLeaves.size()) {
3300 return rewriter.notifyMatchFailure(
3301 op, "expected array-of-pod containment operands to expand to the same number of leaves"
3302 );
3303 }
3304
3305 auto getShapeSource =
3306 [&loc, &rewriter](ArrayType arrTy, Value originalValue, ValueRange convertedValues) {
3307 if (Value carrier = getConvertedPodArrayShapeSource(arrTy, convertedValues)) {
3308 return carrier;
3309 }
3310
3311 if (!convertedValues.empty()) {
3312 return getSingleConvertedValue(convertedValues);
3313 }
3314
3315 return materializeArrayLengthCarrier(originalValue, arrTy, loc, rewriter);
3316 };
3317
3318 Value shapeCarrier = getShapeSource(
3319 compatibleLhsTy, op.getLhs(),
3320 lhsCompatibleConvertedValues.empty() ? adaptor.getLhs()
3321 : ValueRange(lhsCompatibleConvertedValues)
3322 );
3323 Value zero = rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0));
3324 Value trueVal = rewriter.create<arith::ConstantOp>(
3325 loc, IntegerAttr::get(IntegerType::get(rewriter.getContext(), 1), 1)
3326 );
3327
3328 SmallVector<Value> selectedIndices;
3329 selectedIndices.reserve(selectedDims);
3330 for (size_t dim = 0; dim < selectedDims; ++dim) {
3331 Value idx = rewriter.create<NonDetOp>(loc, IndexType::get(rewriter.getContext()));
3332 Value dimVal = rewriter.create<arith::ConstantOp>(
3333 loc, rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3334 );
3335 Value dimLen = rewriter.create<ArrayLengthOp>(loc, shapeCarrier, dimVal);
3336
3337 Value nonNegative = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::sge, idx, zero);
3338 rewriter.create<constrain::EmitEqualityOp>(loc, nonNegative, trueVal);
3339
3340 Value inRange = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::slt, idx, dimLen);
3341 rewriter.create<constrain::EmitEqualityOp>(loc, inRange, trueVal);
3342
3343 selectedIndices.push_back(idx);
3344 }
3345
3346 bool lhsNeedsShapeCheck = rhsArrTy && needsPodArrayShapeCarrier(compatibleLhsTy);
3347 bool rhsNeedsShapeCheck = compatibleRhsArrTy && needsPodArrayShapeCarrier(compatibleRhsArrTy);
3348 if (compatibleRhsArrTy && (lhsNeedsShapeCheck || rhsNeedsShapeCheck)) {
3349 Value rhsShapeSource = getShapeSource(
3350 compatibleRhsArrTy, op.getRhs(),
3351 rhsCompatibleConvertedValues.empty() ? adaptor.getRhs()
3352 : ValueRange(rhsCompatibleConvertedValues)
3353 );
3354 Value selectedShapeSource =
3355 selectedIndices.empty()
3356 ? shapeCarrier
3357 : ArrayAccessOpInterface::genRead(rewriter, loc, shapeCarrier, selectedIndices);
3358 for (size_t dim = 0; dim < rhsRank; ++dim) {
3359 Value dimVal = rewriter.create<arith::ConstantOp>(
3360 loc, rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3361 );
3362 Value lhsLen = rewriter.create<ArrayLengthOp>(loc, selectedShapeSource, dimVal);
3363 Value rhsLen = rewriter.create<ArrayLengthOp>(loc, rhsShapeSource, dimVal);
3364 preserveDiscardableAttrs(
3365 op, rewriter.create<constrain::EmitEqualityOp>(loc, lhsLen, rhsLen)
3366 );
3367 }
3368 }
3369
3370 if (lhsLeaves.empty() && rhsLeaves.empty()) {
3371 if (rhsArrTy) {
3372 Value rhsShapeCarrier = getShapeSource(rhsArrTy, op.getRhs(), adaptor.getRhs());
3373 Value selectedShape =
3374 selectedIndices.empty()
3375 ? shapeCarrier
3376 : rewriter.create<ExtractArrayOp>(
3377 loc, getPodArrayShapeCarrierType(rhsArrTy), shapeCarrier, selectedIndices
3378 );
3379 preserveDiscardableAttrs(
3380 op, rewriter.create<constrain::EmitEqualityOp>(loc, selectedShape, rhsShapeCarrier)
3381 );
3382 }
3383 rewriter.eraseOp(op);
3384 return success();
3385 }
3386
3387 for (auto [lhsLeaf, rhsLeaf] : llvm::zip_equal(lhsLeaves, rhsLeaves)) {
3388 Value selectedLhs = lhsLeaf;
3389 if (auto rhsLeafArrTy = llvm::dyn_cast<ArrayType>(rhsLeaf.getType())) {
3390 if (!selectedIndices.empty()) {
3391 selectedLhs =
3392 rewriter.create<ExtractArrayOp>(loc, rhsLeafArrTy, lhsLeaf, selectedIndices);
3393 }
3394 } else {
3395 selectedLhs =
3396 rewriter.create<ReadArrayOp>(loc, rhsLeaf.getType(), lhsLeaf, selectedIndices);
3397 }
3398 preserveDiscardableAttrs(
3399 op, rewriter.create<constrain::EmitEqualityOp>(loc, selectedLhs, rhsLeaf)
3400 );
3401 }
3402
3403 rewriter.eraseOp(op);
3404 return success();
3405 }
3406};
3407
3414static Value selectArrayLengthShapeSource(
3415 ArrayLengthOp op, ValueRange convertedArrRefs, ConversionPatternRewriter &rewriter
3416) {
3417 if (Value v = getConvertedPodArrayShapeSource(op.getArrRefType(), convertedArrRefs)) {
3418 return v;
3419 }
3420
3421 return materializeArrayLengthCarrier(op.getArrRef(), op.getArrRefType(), op.getLoc(), rewriter);
3422}
3423
3424static Value tryResolveReadPodArrayShapeSource(
3425 ReadPodOp readOp, ArrayType arrTy, const VirtualPodValueMap &virtualPods, Location loc,
3426 RewriterBase &rewriter
3427);
3428
3430class RejectRaggedNestedLeafArrayLengthOp : public OpConversionPattern<ArrayLengthOp> {
3431public:
3432 using OpConversionPattern<ArrayLengthOp>::OpConversionPattern;
3433
3434 static bool legal(ArrayLengthOp op) {
3435 return getTaggedRaggedNestedLeafKind(op.getArrRef()).empty();
3436 }
3437
3438 LogicalResult
3439 matchAndRewrite(ArrayLengthOp op, OneToNOpAdaptor, ConversionPatternRewriter &) const override {
3440 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getArrRef());
3441 if (raggedKind.empty()) {
3442 return failure();
3443 }
3444 return op.emitOpError() << "cannot lower nested " << raggedKind
3445 << " array leaf length after reading an array-of-POD element without "
3446 "per-element shape witnesses";
3447 }
3448};
3449
3451class SplitPodArrayLengthOp : public OpConversionPattern<ArrayLengthOp> {
3452public:
3453 using OpConversionPattern<ArrayLengthOp>::OpConversionPattern;
3454
3455 static bool legal(ArrayLengthOp op) {
3456 return RejectRaggedNestedLeafArrayLengthOp::legal(op) &&
3457 (!splittablePodArray(op.getArrRefType()) || op->hasAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR));
3458 }
3459
3460 LogicalResult matchAndRewrite(
3461 ArrayLengthOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3462 ) const override {
3463 if (legal(op)) {
3464 return failure();
3465 }
3466 if (shouldDeferPodArrayLengthToStep3(op)) {
3467 auto deferred = rewriter.create<ArrayLengthOp>(
3468 op.getLoc(), op.getArrRef(), getSingleConvertedValue(adaptor.getDim())
3469 );
3470 preserveDiscardableAttrs(op, deferred);
3471 deferred->setAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR, UnitAttr::get(op.getContext()));
3472 rewriter.replaceOp(op, deferred.getResult());
3473 return success();
3474 }
3475 Value arrRef = selectArrayLengthShapeSource(op, adaptor.getArrRef(), rewriter);
3476 preserveDiscardableAttrs(
3477 op, rewriter.replaceOpWithNewOp<ArrayLengthOp>(
3478 op, arrRef, getSingleConvertedValue(adaptor.getDim())
3479 )
3480 );
3481 return success();
3482 }
3483};
3484
3490struct DeferredPodArrayBacking {
3491 SmallVector<Value> leafArrays;
3492 Value shapeCarrier;
3493};
3494
3495using DeferredPodArrayBackingMap = DenseMap<Value, DeferredPodArrayBacking>;
3496
3498static void setDeferredPodArrayBackingInsertionPoint(ReadPodOp readOp, OpBuilder &bldr) {
3499 if (Operation *loopOp = llzk::pod::detail::findNearestLoopCarriedPodAccess(readOp)) {
3500 bldr.setInsertionPoint(loopOp);
3501 } else {
3502 bldr.setInsertionPointAfter(readOp);
3503 }
3504}
3505
3507static DeferredPodArrayBacking &materializeDeferredPodArrayBacking(
3508 ReadPodOp readOp, ArrayType arrTy, ArrayRef<Type> splitTypes,
3509 DeferredPodArrayBackingMap &deferredPodArrays, Location loc, OpBuilder &bldr,
3510 bool requireLeafArrays, bool requireShapeCarrier
3511) {
3512 auto [it, inserted] = deferredPodArrays.try_emplace(readOp.getResult());
3513 DeferredPodArrayBacking &backing = it->second;
3514
3515 bool needsShapeCarrier = requireShapeCarrier && needsPodArrayShapeCarrier(arrTy);
3516 bool missingLeafArrays = requireLeafArrays && backing.leafArrays.empty();
3517 bool missingShapeCarrier = needsShapeCarrier && !backing.shapeCarrier;
3518 if (missingLeafArrays || missingShapeCarrier) {
3519 OpBuilder::InsertionGuard guard(bldr);
3520 setDeferredPodArrayBackingInsertionPoint(readOp, bldr);
3521
3522 if (missingLeafArrays) {
3523 backing.leafArrays.reserve(splitTypes.size());
3524 for (Type splitType : splitTypes) {
3525 backing.leafArrays.push_back(
3526 createWritableArrayValue(bldr, loc, llvm::cast<ArrayType>(splitType))
3527 );
3528 }
3529 }
3530
3531 if (missingShapeCarrier) {
3532 backing.shapeCarrier = materializeArrayLengthCarrier(readOp.getResult(), arrTy, loc, bldr);
3533 }
3534 }
3535
3536 if (inserted || requireLeafArrays) {
3537 assert(
3538 backing.leafArrays.size() == splitTypes.size() &&
3539 "cached split POD arrays must match the rewritten read arity"
3540 );
3541 }
3542 return backing;
3543}
3544
3546struct Step3Resolver {
3547 VirtualPodValueMap virtualPods;
3548 CompatiblePodLeafMaterializationMap materializedLeaves;
3549 DeferredPodArrayBackingMap deferredPodArrays;
3550
3551 void rehydrateVirtualPodPlaceholders(ModuleOp modOp);
3552 void addPreConversionPatterns(RewritePatternSet &patterns);
3553 void addConversionPatterns(
3554 RewritePatternSet &patterns, SymbolTableCollection &symTables,
3555 const MemberReplacementMap &memberRepMap
3556 );
3557 void addLateResolutionPatterns(RewritePatternSet &patterns);
3558 void addPostConversionPatterns(RewritePatternSet &patterns);
3559 void configureLateVirtualPodLegality(ConversionTarget &target) const;
3560 bool hasResolvableLateVirtualPodOps(ModuleOp modOp) const;
3561 void materializeRemainingVirtualPods(ModuleOp modOp);
3562};
3563
3565class ResolvePodReadBackedArrayLengthOp final : public OpConversionPattern<ArrayLengthOp> {
3566 Step3Resolver &resolver;
3567
3568public:
3569 ResolvePodReadBackedArrayLengthOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
3570 : OpConversionPattern<ArrayLengthOp>(ctx), resolver(step3Resolver) {}
3571
3572 static bool legal(ArrayLengthOp op) { return !op->hasAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR); }
3573
3574 LogicalResult
3575 matchAndRewrite(ArrayLengthOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
3576 ArrayType arrTy = splittablePodArray(op.getArrRefType());
3577 if (!arrTy || !shouldDeferPodArrayLengthToStep3(op)) {
3578 return failure();
3579 }
3580
3581 Value shapeSource;
3582 ReadPodOp readOp = getReadPodBacking(op.getArrRef());
3583 assert(readOp && "deferred POD-backed array.len must still trace to pod.read");
3584 shapeSource = tryResolveReadPodArrayShapeSource(
3585 readOp, arrTy, resolver.virtualPods, op.getLoc(), rewriter
3586 );
3587
3588 if (!shapeSource) {
3589 if (needsPodArrayShapeCarrier(arrTy) && isFreshUnwrittenPodRead(readOp)) {
3590 shapeSource =
3591 materializeDeferredPodArrayBacking(
3592 readOp, arrTy, /*splitTypes=*/ArrayRef<Type> {}, resolver.deferredPodArrays,
3593 op.getLoc(), rewriter, /*requireLeafArrays=*/false,
3594 /*requireShapeCarrier=*/true
3595 )
3596 .shapeCarrier;
3597 } else if (auto it = resolver.deferredPodArrays.find(readOp.getResult());
3598 it != resolver.deferredPodArrays.end() && it->second.shapeCarrier) {
3599 shapeSource = castValueToTypeIfNeeded(
3600 rewriter, op.getLoc(), it->second.shapeCarrier, getPodArrayShapeCarrierType(arrTy)
3601 );
3602 }
3603 }
3604
3605 if (!shapeSource) {
3606 shapeSource = materializeArrayLengthCarrier(op.getArrRef(), arrTy, op.getLoc(), rewriter);
3607 }
3608
3609 preserveDiscardableAttrsExcept(
3610 op, rewriter.replaceOpWithNewOp<ArrayLengthOp>(op, shapeSource, op.getDim()),
3611 DEFERRED_POD_ARRAY_LENGTH_ATTR
3612 );
3613 return success();
3614 }
3615};
3616
3618class SplitPodArrayExtractArrayOp : public OpConversionPattern<ExtractArrayOp> {
3619public:
3620 using OpConversionPattern<ExtractArrayOp>::OpConversionPattern;
3621
3622 static bool legal(ExtractArrayOp op) { return !splittablePodArray(op.getResult().getType()); }
3623
3624 LogicalResult matchAndRewrite(
3625 ExtractArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3626 ) const override {
3627 if (legal(op)) {
3628 return failure();
3629 }
3630
3631 SmallVector<Type> splitResultTypes;
3632 splitPodArrayTypeTo(op.getResult().getType(), splitResultTypes);
3633 if (splitResultTypes.empty()) {
3634 ArrayType resultTy = llvm::cast<ArrayType>(op.getResult().getType());
3635 preserveDiscardableAttrs(
3636 op, rewriter.replaceOpWithNewOp<ExtractArrayOp>(
3637 op, getPodArrayShapeCarrierType(resultTy),
3638 getSingleConvertedValue(adaptor.getArrRef()),
3639 flattenConvertedValues(adaptor.getIndices())
3640 )
3641 );
3642 return success();
3643 }
3644
3645 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
3646 SmallVector<Value> replacements;
3647 ArrayType resultTy = llvm::cast<ArrayType>(op.getResult().getType());
3648 replacements.reserve(splitResultTypes.size() + (needsPodArrayShapeCarrier(resultTy) ? 1 : 0));
3649 auto splitArrRefs = adaptor.getArrRef().take_front(splitResultTypes.size());
3650 for (auto [splitArrRange, splitResultType] : llvm::zip_equal(splitArrRefs, splitResultTypes)) {
3651 replacements.push_back(preserveDiscardableAttrs(
3652 op, rewriter.create<ExtractArrayOp>(
3653 op.getLoc(), llvm::cast<ArrayType>(splitResultType),
3654 getSingleConvertedValue(splitArrRange), indices
3655 )
3656 ));
3657 }
3658 if (needsPodArrayShapeCarrier(resultTy)) {
3659 Value shapeCarrier = materializeExtractedPodArrayShapeCarrier(
3660 op, resultTy, op.getArrRef(), adaptor.getArrRef(), indices, rewriter
3661 );
3662 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
3663 replacements.push_back(shapeCarrier);
3664 }
3665
3666 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
3667 return success();
3668 }
3669};
3670
3672class SplitPodArrayInsertArrayOp : public OpConversionPattern<InsertArrayOp> {
3673public:
3674 using OpConversionPattern<InsertArrayOp>::OpConversionPattern;
3675
3676 static bool legal(InsertArrayOp op) { return !splittablePodArray(op.getRvalue().getType()); }
3677
3678 LogicalResult matchAndRewrite(
3679 InsertArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3680 ) const override {
3681 if (legal(op)) {
3682 return failure();
3683 }
3684
3685 if (hasZeroLeafPodArraySplit(llvm::cast<ArrayType>(op.getRvalue().getType()))) {
3686 preserveDiscardableAttrs(
3687 op, rewriter.create<InsertArrayOp>(
3688 op.getLoc(), getSingleConvertedValue(adaptor.getArrRef()),
3689 flattenConvertedValues(adaptor.getIndices()),
3690 getSingleConvertedValue(adaptor.getRvalue())
3691 )
3692 );
3693 rewriter.eraseOp(op);
3694 return success();
3695 }
3696
3697 ArrayType destArrTy = llvm::cast<ArrayType>(op.getArrRef().getType());
3698 ArrayType rvalueTy = llvm::cast<ArrayType>(op.getRvalue().getType());
3699 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
3700 size_t leafCount = getSplitPodArrayLeafCount(rvalueTy);
3701 auto splitArrRefs = adaptor.getArrRef().take_front(leafCount);
3702 auto splitRvalues = adaptor.getRvalue().take_front(leafCount);
3703 for (auto [splitArrRange, splitRvalueRange] : llvm::zip_equal(splitArrRefs, splitRvalues)) {
3704 preserveDiscardableAttrs(
3705 op, rewriter.create<InsertArrayOp>(
3706 op.getLoc(), getSingleConvertedValue(splitArrRange), indices,
3707 getSingleConvertedValue(splitRvalueRange)
3708 )
3709 );
3710 }
3711 if (needsPodArrayShapeCarrier(destArrTy)) {
3712 Value destCarrier = getConvertedPodArrayShapeCarrierIfPresent(destArrTy, adaptor.getArrRef());
3713 if (!destCarrier) {
3714 return rewriter.notifyMatchFailure(
3715 op, "expected converted destination shape carrier for array-of-pod insert"
3716 );
3717 }
3718
3719 Value rvalueCarrier =
3720 getConvertedPodArrayShapeCarrierIfPresent(rvalueTy, adaptor.getRvalue());
3721 if (!rvalueCarrier) {
3722 rvalueCarrier =
3723 materializeArrayLengthCarrier(op.getRvalue(), rvalueTy, op.getLoc(), rewriter);
3724 }
3725
3726 preserveDiscardableAttrs(
3727 op, rewriter.create<InsertArrayOp>(op.getLoc(), destCarrier, indices, rvalueCarrier)
3728 );
3729 }
3730
3731 rewriter.eraseOp(op);
3732 return success();
3733 }
3734};
3735
3737class SplitPodArrayInMemberWriteOp : public OpConversionPattern<MemberWriteOp> {
3738 SymbolTableCollection &tables;
3739 const MemberReplacementMap &repMapRef;
3740
3741public:
3742 SplitPodArrayInMemberWriteOp(
3743 const TypeConverter &converter, MLIRContext *ctx, SymbolTableCollection &symTables,
3744 const MemberReplacementMap &memberRepMap
3745 )
3746 : OpConversionPattern<MemberWriteOp>(converter, ctx), tables(symTables),
3747 repMapRef(memberRepMap) {}
3748
3749 static bool legal(MemberWriteOp op) { return !splittablePodArray(op.getVal().getType()); }
3750
3751 LogicalResult matchAndRewrite(
3752 MemberWriteOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3753 ) const override {
3754 if (legal(op)) {
3755 return failure();
3756 }
3757 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
3758 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
3759 assert(succeeded(tgtStructDef));
3760
3761 const LocalMemberReplacementMap &idToMember =
3762 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
3763 ArrayType arrTy = llvm::cast<ArrayType>(op.getVal().getType());
3764 SmallVector<RecordChain> splitIds;
3765 SmallVector<Type> splitTypes;
3766 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
3767 if (splitTypes.empty()) {
3768 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3769 preserveDiscardableAttrs(
3770 op,
3771 rewriter.create<MemberWriteOp>(
3772 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3773 FlatSymbolRefAttr::get(carrierMember.first), getSingleConvertedValue(adaptor.getVal())
3774 )
3775 );
3776 rewriter.eraseOp(op);
3777 return success();
3778 }
3779
3780 auto splitVals = adaptor.getVal().take_front(splitIds.size());
3781 for (auto [id, splitValRange] : llvm::zip_equal(splitIds, splitVals)) {
3782 const MemberInfo &newMember = idToMember.at(id);
3783 preserveDiscardableAttrs(
3784 op, rewriter.create<MemberWriteOp>(
3785 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3786 FlatSymbolRefAttr::get(newMember.first), getSingleConvertedValue(splitValRange)
3787 )
3788 );
3789 }
3790 if (needsPodArrayShapeCarrier(arrTy)) {
3791 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, adaptor.getVal());
3792 if (!carrier) {
3793 carrier = materializeArrayLengthCarrier(op.getVal(), arrTy, op.getLoc(), rewriter);
3794 }
3795 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3796 preserveDiscardableAttrs(
3797 op, rewriter.create<MemberWriteOp>(
3798 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3799 FlatSymbolRefAttr::get(carrierMember.first),
3800 castValueToTypeIfNeeded(rewriter, op.getLoc(), carrier, carrierMember.second)
3801 )
3802 );
3803 }
3804 rewriter.eraseOp(op);
3805 return success();
3806 }
3807};
3808
3810class SplitPodArrayInMemberReadOp : public OpConversionPattern<MemberReadOp> {
3811 SymbolTableCollection &tables;
3812 const MemberReplacementMap &repMapRef;
3813
3814public:
3815 SplitPodArrayInMemberReadOp(
3816 const TypeConverter &converter, MLIRContext *ctx, SymbolTableCollection &symTables,
3817 const MemberReplacementMap &memberRepMap
3818 )
3819 : OpConversionPattern<MemberReadOp>(converter, ctx), tables(symTables),
3820 repMapRef(memberRepMap) {}
3821
3822 static bool legal(MemberReadOp op) { return !splittablePodArray(op.getResult().getType()); }
3823
3824 LogicalResult matchAndRewrite(
3825 MemberReadOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3826 ) const override {
3827 if (legal(op)) {
3828 return failure();
3829 }
3830 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
3831 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
3832 assert(succeeded(tgtStructDef));
3833
3834 const LocalMemberReplacementMap &idToMember =
3835 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
3836 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
3837 SmallVector<RecordChain> splitIds;
3838 SmallVector<Type> splitTypes;
3839 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
3840 SmallVector<Value> mapOperands;
3841 std::optional<int32_t> numDimsPerMap;
3842 auto mapOperandsOld = adaptor.getMapOperands();
3843 if (!mapOperandsOld.empty()) {
3844 assert(
3845 mapOperandsOld.size() == 1 &&
3846 "member.readm should have at most one affine-map operand group"
3847 );
3848 mapOperands = flattenConvertedValues(mapOperandsOld.front());
3849
3850 ArrayRef<int32_t> numDimsPerMapOld = op.getNumDimsPerMap();
3851 if (!numDimsPerMapOld.empty()) {
3852 assert(
3853 numDimsPerMapOld.size() == 1 &&
3854 "member.readm should have one numDims entry per affine-map group"
3855 );
3856 numDimsPerMap = numDimsPerMapOld.front();
3857 }
3858 }
3859 if (splitTypes.empty()) {
3860 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3861 Value carrierRead = preserveDiscardableAttrs(
3862 op,
3863 rewriter.create<MemberReadOp>(
3864 op.getLoc(), carrierMember.second, getSingleConvertedValue(adaptor.getComponent()),
3865 carrierMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3866 )
3867 );
3868 rewriter.replaceOpWithMultiple(op, {ValueRange {carrierRead}});
3869 return success();
3870 }
3871 SmallVector<Value> replacements;
3872 replacements.reserve(splitIds.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
3873 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
3874 const MemberInfo &newMember = idToMember.at(id);
3875 replacements.push_back(preserveDiscardableAttrs(
3876 op, rewriter.create<MemberReadOp>(
3877 op.getLoc(), splitType, getSingleConvertedValue(adaptor.getComponent()),
3878 newMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3879 )
3880 ));
3881 }
3882 if (needsPodArrayShapeCarrier(arrTy)) {
3883 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3884 replacements.push_back(preserveDiscardableAttrs(
3885 op,
3886 rewriter.create<MemberReadOp>(
3887 op.getLoc(), carrierMember.second, getSingleConvertedValue(adaptor.getComponent()),
3888 carrierMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3889 )
3890 ));
3891 }
3892 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
3893 return success();
3894 }
3895};
3896
3900static LogicalResult rejectUnsupportedPodArrayUnifiableCasts(ModuleOp modOp) {
3901 WalkResult result = modOp.walk([](UnifiableCastOp op) {
3902 ArrayType inputArrTy = splittablePodArray(op.getInput().getType());
3903 ArrayType resultArrTy = splittablePodArray(op.getType());
3904 if (!inputArrTy && !resultArrTy) {
3905 return WalkResult::advance();
3906 }
3907 if (!inputArrTy && resultArrTy) {
3908 op.emitOpError()
3909 .append(
3910 "cannot lower a generic input to a split array-of-POD result without rewriting "
3911 "the surrounding function/template signature"
3912 )
3913 .report();
3914 return WalkResult::interrupt();
3915 }
3916 if (inputArrTy && !resultArrTy) {
3917 op.emitOpError()
3918 .append(
3919 "cannot lower a split array-of-POD input to a non-array result without "
3920 "materializing the aggregate or rewriting the surrounding function/template "
3921 "signature"
3922 )
3923 .report();
3924 return WalkResult::interrupt();
3925 }
3926 return WalkResult::advance();
3927 });
3928 return failure(result.wasInterrupted());
3929}
3930
3932static LogicalResult
3933step2(ModuleOp modOp, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap) {
3934 if (failed(rejectUnsupportedPodArrayUnifiableCasts(modOp))) {
3935 return failure();
3936 }
3937
3938 MLIRContext *ctx = modOp.getContext();
3939 PodArrayTypeConverter typeConverter;
3940 CompatiblePodLeafMaterializationMap materializedLeaves;
3941
3942 RewritePatternSet patterns(ctx);
3943 patterns.add<
3944 SplitPodArrayNonDetOp, SplitPodArrayCreateArrayOp, SplitPodArrayReadArrayOp,
3945 SplitPodArrayWriteArrayOp, SplitPodArrayExtractArrayOp, SplitPodArrayInsertArrayOp,
3946 SplitPodArrayInFuncDefOp, SplitPodArrayInUnifiableCastOp, SplitPodArrayInReturnOp,
3947 SplitPodArrayInCallOp, RejectRaggedNestedLeafArrayLengthOp, SplitPodArrayLengthOp>(
3948 typeConverter, ctx
3949 );
3950 patterns.add<SplitPodArrayInEmitEqualityOp, SplitPodArrayInEmitContainmentOp>(
3951 typeConverter, ctx, materializedLeaves
3952 );
3953 patterns.add<SplitPodArrayInMemberWriteOp, SplitPodArrayInMemberReadOp>(
3954 typeConverter, ctx, symTables, memberRepMap
3955 );
3956
3957 ConversionTarget target(*ctx);
3958 baseTargetSetup(target);
3959 target.addLegalOp<UnrealizedConversionCastOp>();
3960 target.addDynamicallyLegalOp<NonDetOp>(SplitPodArrayNonDetOp::legal);
3961 target.addDynamicallyLegalOp<CreateArrayOp>(SplitPodArrayCreateArrayOp::legal);
3962 target.addDynamicallyLegalOp<ReadArrayOp>(SplitPodArrayReadArrayOp::legal);
3963 target.addDynamicallyLegalOp<WriteArrayOp>(SplitPodArrayWriteArrayOp::legal);
3964 target.addDynamicallyLegalOp<ExtractArrayOp>(SplitPodArrayExtractArrayOp::legal);
3965 target.addDynamicallyLegalOp<InsertArrayOp>(SplitPodArrayInsertArrayOp::legal);
3966 target.addDynamicallyLegalOp<FuncDefOp>(SplitPodArrayInFuncDefOp::legal);
3967 target.addDynamicallyLegalOp<UnifiableCastOp>(SplitPodArrayInUnifiableCastOp::legal);
3968 target.addDynamicallyLegalOp<ReturnOp>(SplitPodArrayInReturnOp::legal);
3969 target.addDynamicallyLegalOp<CallOp>(SplitPodArrayInCallOp::legal);
3970 target.addDynamicallyLegalOp<constrain::EmitEqualityOp>(SplitPodArrayInEmitEqualityOp::legal);
3971 target.addDynamicallyLegalOp<constrain::EmitContainmentOp>(
3972 SplitPodArrayInEmitContainmentOp::legal
3973 );
3974 target.addDynamicallyLegalOp<ArrayLengthOp>(SplitPodArrayLengthOp::legal);
3975 target.addDynamicallyLegalOp<MemberWriteOp>(SplitPodArrayInMemberWriteOp::legal);
3976 target.addDynamicallyLegalOp<MemberReadOp>(SplitPodArrayInMemberReadOp::legal);
3977
3978 mlir::scf::populateSCFStructuralTypeConversionsAndLegality(typeConverter, patterns, target);
3979
3980 LLVM_DEBUG(llvm::dbgs() << "Begin step 2: split arrays with POD element type\n";);
3981 return applyPartialConversion(modOp, target, std::move(patterns));
3982}
3983
3985class SplitInitFromNewPodOp : public OpConversionPattern<NewPodOp> {
3986public:
3987 using OpConversionPattern<NewPodOp>::OpConversionPattern;
3988
3989 static bool legal(NewPodOp op) { return op.getInitialValues().empty(); }
3990
3991 LogicalResult match(NewPodOp op) const override { return failure(legal(op)); }
3992
3993 void rewrite(NewPodOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
3994 // Generate an individual write for each initialization
3995 rewriter.setInsertionPointAfter(op);
3996 Location loc = op.getLoc();
3997 for (auto [name, init] :
3998 llvm::zip_equal(adaptor.getInitializedRecords(), adaptor.getInitialValues())) {
3999 // Create the write
4000 rewriter.create<WritePodOp>(loc, op.getResult(), llvm::cast<StringAttr>(name), init);
4001 }
4002 // Remove initializations from `op`
4003 rewriter.modifyOpInPlace(op, [&op]() {
4004 op.getInitialValuesMutable().clear();
4005 op.setInitializedRecordsAttr(ArrayAttr::get(op.getContext(), {})); // DefaultValuedAttr:{}
4006 });
4007 }
4008};
4009
4016class SplitPodElementCreateArrayOp : public OpConversionPattern<CreateArrayOp> {
4017 const Step3Resolver &resolver;
4018
4019public:
4020 SplitPodElementCreateArrayOp(MLIRContext *ctx, const Step3Resolver &step3Resolver)
4021 : OpConversionPattern<CreateArrayOp>(ctx), resolver(step3Resolver) {}
4022
4023 static bool legal(CreateArrayOp op) {
4024 return !llvm::any_of(op.getElements().getTypes(), [](Type type) {
4025 return splittablePod(type) || llvm::isa<ArrayType>(type);
4026 });
4027 }
4028
4029 LogicalResult match(CreateArrayOp op) const override { return failure(legal(op)); }
4030
4031 void
4032 rewrite(CreateArrayOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
4033 SmallVector<Value> leafElements;
4034 leafElements.reserve(adaptor.getElements().size());
4035
4036 Type leafType;
4037 for (Value element : adaptor.getElements()) {
4038 SmallVector<Value> flattenedValues;
4039 if (splittablePod(element.getType())) {
4040 processInputOperand(
4041 op.getLoc(), element, flattenedValues, rewriter, op.getOperation(),
4042 &resolver.virtualPods
4043 );
4044 } else {
4045 flattenedValues.push_back(element);
4046 }
4047
4048 assert(
4049 flattenedValues.size() == 1 &&
4050 "array.new elements should already have been split to a single flattened leaf"
4051 );
4052 if (!leafType) {
4053 leafType = flattenedValues.front().getType();
4054 } else {
4055 assert(
4056 leafType == flattenedValues.front().getType() && "array.new elements must stay uniform"
4057 );
4058 }
4059 leafElements.push_back(flattenedValues.front());
4060 }
4061
4062 size_t leafRank = 0;
4063 if (auto leafArrTy = llvm::dyn_cast_if_present<ArrayType>(leafType)) {
4064 leafRank = leafArrTy.getDimensionSizes().size();
4065 }
4066 ArrayType arrTy = op.getType();
4067 assert(
4068 arrTy.getDimensionSizes().size() >= leafRank && "flattened leaf rank exceeds array rank"
4069 );
4070 size_t outerRank = arrTy.getDimensionSizes().size() - leafRank;
4071 assert(outerRank > 0 && "array.new elements must populate at least one outer array dimension");
4072
4073 ArrayType outerIndexTy =
4074 ArrayType::get(arrTy.getElementType(), arrTy.getDimensionSizes().take_front(outerRank));
4075 auto elementIndices = outerIndexTy.getSubelementIndices();
4076 assert(
4077 elementIndices && "array.new with explicit POD elements requires static outer dimensions"
4078 );
4079 assert(
4080 elementIndices->size() == leafElements.size() &&
4081 "array.new element count must match the outer array cardinality"
4082 );
4083
4084 Value rebuiltArray = createWritableArrayValue(rewriter, op.getLoc(), arrTy);
4085 preserveDiscardableAttrs(op, rebuiltArray.getDefiningOp());
4086 for (auto [index, leafValue] : llvm::zip_equal(*elementIndices, leafElements)) {
4087 genArrayWrite(rewriter, op.getLoc(), rebuiltArray, index, leafValue);
4088 }
4089 rewriter.replaceOp(op, rebuiltArray);
4090 }
4091};
4092
4100class SplitPodInFuncDefOp : public OpConversionPattern<FuncDefOp> {
4101 Step3Resolver &resolver;
4102
4103public:
4104 SplitPodInFuncDefOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4105 : OpConversionPattern<FuncDefOp>(ctx), resolver(step3Resolver) {}
4106
4107 inline static bool legal(FuncDefOp op) {
4108 return !containsSplittablePodType(op.getArgumentTypes()) &&
4109 !containsSplittablePodType(op.getResultTypes());
4110 }
4111
4112 LogicalResult match(FuncDefOp op) const override { return failure(legal(op)); }
4113
4114 void rewrite(FuncDefOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
4115 // Update in/out types of the function to replace pods with scalars
4116 class Impl : public FunctionTypeConverter {
4117 SmallVector<size_t> originalInputIdxToSize, originalResultIdxToSize;
4118 SplitFunctionNameInfo inputNameInfo;
4119 SplitFunctionNameInfo resultNameInfo;
4120 Step3Resolver &resolver;
4121
4122 protected:
4123 SmallVector<Type> convertInputs(ArrayRef<Type> origTypes) override {
4124 return splitPodType(origTypes, &originalInputIdxToSize);
4125 }
4126 SmallVector<Type> convertResults(ArrayRef<Type> origTypes) override {
4127 return splitPodType(origTypes, &originalResultIdxToSize);
4128 }
4129 ArrayAttr convertInputAttrs(ArrayAttr origAttrs, SmallVector<Type> newTypes) override {
4131 origAttrs, originalInputIdxToSize, newTypes, ARG_NAME_ATTR_NAME,
4132 inputNameInfo.originalNames, inputNameInfo.existingNames,
4133 inputNameInfo.splitNameSuffixes
4134 );
4135 }
4136 ArrayAttr convertResultAttrs(ArrayAttr origAttrs, SmallVector<Type> newTypes) override {
4138 origAttrs, originalResultIdxToSize, newTypes, RES_NAME_ATTR_NAME,
4139 resultNameInfo.originalNames, resultNameInfo.existingNames,
4140 resultNameInfo.splitNameSuffixes
4141 );
4142 }
4143
4148 void processBlockArgs(Block &entryBlock, RewriterBase &rewriter) override {
4149 OpBuilder::InsertionGuard guard(rewriter);
4150 rewriter.setInsertionPointToStart(&entryBlock);
4151
4152 for (unsigned i = 0; i < entryBlock.getNumArguments();) {
4153 Value oldV = entryBlock.getArgument(i);
4154 if (PodType pt = splittablePod(oldV.getType())) {
4155 Location loc = oldV.getLoc();
4156 VirtualPodLeafMap leafValues;
4157 SmallVector<StringAttr> recordChain;
4158 unsigned nextArgIdx = i + 1;
4159 forEachPodLeaf(pt, recordChain, [&](const RecordChain &id, Type leafType) {
4160 BlockArgument newArg = entryBlock.insertArgument(nextArgIdx, leafType, loc);
4161 leafValues[id] = newArg;
4162 ++nextArgIdx;
4163 });
4164
4165 Value virtualPod = createVirtualPodPlaceholder(rewriter, loc, pt, leafValues);
4166 rewriter.replaceAllUsesWith(oldV, virtualPod);
4167 entryBlock.eraseArgument(i);
4168
4169 i += leafValues.size();
4170 resolver.virtualPods[virtualPod] = std::move(leafValues);
4171 } else {
4172 ++i;
4173 }
4174 }
4175 }
4176
4177 public:
4178 Impl(FuncDefOp op, Step3Resolver &step3Resolver) : resolver(step3Resolver) {
4179 inputNameInfo = collectSplitFunctionNameInfo(op.getArgumentTypes(), [&op](unsigned i) {
4180 return op.getArgNameAttr(i);
4181 }, getSplitRecordNameSuffixes);
4182 resultNameInfo = collectSplitFunctionNameInfo(
4183 op.getResultTypes(), [resultAttrs = op.getAllResultAttrs()](unsigned i) {
4184 return getAttrAtIndexWithName(resultAttrs, i, RES_NAME_ATTR_NAME);
4185 }, getSplitRecordNameSuffixes
4186 );
4187 }
4188 };
4189 Impl(op, resolver).convert(op, rewriter);
4190 }
4191};
4192
4198class SplitPodInReturnOp : public OpConversionPattern<ReturnOp> {
4199 const Step3Resolver &resolver;
4200
4201public:
4202 SplitPodInReturnOp(MLIRContext *ctx, const Step3Resolver &step3Resolver)
4203 : OpConversionPattern<ReturnOp>(ctx), resolver(step3Resolver) {}
4204
4205 inline static bool legal(ReturnOp op) {
4206 return !containsSplittablePodType(op.getOperands().getTypes());
4207 }
4208
4209 LogicalResult match(ReturnOp op) const override { return failure(legal(op)); }
4210
4211 void rewrite(ReturnOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
4212 processInputOperands(
4213 adaptor.getOperands(), op.getOperandsMutable(), op, rewriter, &resolver.virtualPods
4214 );
4215 }
4216};
4217
4219static CallOp newCallOpWithSplitResults(
4220 CallOp oldCall, CallOp::Adaptor adaptor, ConversionPatternRewriter &rewriter,
4221 Step3Resolver &resolver
4222) {
4223 OpBuilder::InsertionGuard guard(rewriter);
4224 rewriter.setInsertionPointAfter(oldCall);
4225
4226 Operation::result_range oldResults = oldCall.getResults();
4228 oldCall.getLoc(), splitPodType(oldResults.getTypes()), oldCall, adaptor.getMapOperands(),
4229 adaptor.getArgOperands(), rewriter
4230 );
4231
4232 auto newResults = newCall.getResults().begin();
4233 for (Value oldVal : oldResults) {
4234 if (PodType pt = splittablePod(oldVal.getType())) {
4235 Location loc = oldVal.getLoc();
4236 VirtualPodLeafMap leafValues;
4237 SmallVector<StringAttr> recordChain;
4238 forEachPodLeaf(pt, recordChain, [&leafValues, &newResults](const RecordChain &id, Type) {
4239 leafValues[id] = *newResults;
4240 ++newResults;
4241 });
4242 Value virtualPod = createVirtualPodPlaceholder(rewriter, loc, pt, leafValues);
4243 resolver.virtualPods[virtualPod] = std::move(leafValues);
4244 rewriter.replaceAllUsesWith(oldVal, virtualPod);
4245 } else {
4246 rewriter.replaceAllUsesWith(oldVal, *newResults);
4247 newResults++;
4248 }
4249 }
4250 // erase the original CallOp
4251 rewriter.eraseOp(oldCall);
4252
4253 return newCall;
4254}
4255
4262class SplitPodInCallOp : public OpConversionPattern<CallOp> {
4263 Step3Resolver &resolver;
4264
4265public:
4266 SplitPodInCallOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4267 : OpConversionPattern<CallOp>(ctx), resolver(step3Resolver) {}
4268
4269 inline static bool legal(CallOp op) {
4270 return !containsSplittablePodType(op.getArgOperands().getTypes()) &&
4271 !containsSplittablePodType(op.getResultTypes());
4272 }
4273
4274 LogicalResult match(CallOp op) const override { return failure(legal(op)); }
4275
4276 void rewrite(CallOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
4277 // Create new CallOp with split results first so, then process its inputs to split types
4278 CallOp newCall = newCallOpWithSplitResults(op, adaptor, rewriter, resolver);
4279 processInputOperands(
4280 newCall.getArgOperands(), newCall.getArgOperandsMutable(), newCall, rewriter,
4281 &resolver.virtualPods
4282 );
4283 }
4284};
4285
4287class SplitPodInMemberWriteOp : public OpConversionPattern<MemberWriteOp> {
4288 SymbolTableCollection &tables;
4289 const MemberReplacementMap &repMapRef;
4290 const Step3Resolver &resolver;
4291
4292public:
4293 SplitPodInMemberWriteOp(
4294 MLIRContext *ctx, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap,
4295 const Step3Resolver &step3Resolver
4296 )
4297 : OpConversionPattern<MemberWriteOp>(ctx), tables(symTables), repMapRef(memberRepMap),
4298 resolver(step3Resolver) {}
4299
4300 static bool legal(MemberWriteOp op) { return !containsSplittablePodType(op.getVal().getType()); }
4301
4302 LogicalResult match(MemberWriteOp op) const override { return failure(legal(op)); }
4303
4304 void
4305 rewrite(MemberWriteOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
4306 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
4307 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
4308 assert(succeeded(tgtStructDef));
4309
4310 const LocalMemberReplacementMap &idToMember =
4311 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
4312 const VirtualPodLeafMap *virtualLeafValues =
4313 !hasEarlierWriteToPod(op.getOperation(), op.getVal())
4314 ? lookupVirtualPodLeafMap(op.getVal(), resolver.virtualPods)
4315 : nullptr;
4316
4317 for (const auto &[id, newMember] : idToMember) {
4318 Value scalarValue = virtualLeafValues
4319 ? virtualLeafValues->at(id)
4320 : genReadAlongPath(rewriter, op.getLoc(), op.getVal(), id);
4321 preserveDiscardableAttrs(
4322 op, rewriter.create<MemberWriteOp>(
4323 op.getLoc(), adaptor.getComponent(), FlatSymbolRefAttr::get(newMember.first),
4324 scalarValue
4325 )
4326 );
4327 }
4328 rewriter.eraseOp(op);
4329 }
4330};
4331
4333class SplitPodInMemberReadOp : public OpConversionPattern<MemberReadOp> {
4334 SymbolTableCollection &tables;
4335 const MemberReplacementMap &repMapRef;
4336 Step3Resolver &resolver;
4337
4338public:
4339 SplitPodInMemberReadOp(
4340 MLIRContext *ctx, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap,
4341 Step3Resolver &step3Resolver
4342 )
4343 : OpConversionPattern<MemberReadOp>(ctx), tables(symTables), repMapRef(memberRepMap),
4344 resolver(step3Resolver) {}
4345
4346 static bool legal(MemberReadOp op) {
4347 return !containsSplittablePodType(op.getResult().getType());
4348 }
4349
4350 LogicalResult match(MemberReadOp op) const override { return failure(legal(op)); }
4351
4352 void
4353 rewrite(MemberReadOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override {
4354 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
4355 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
4356 assert(succeeded(tgtStructDef));
4357
4358 const LocalMemberReplacementMap &idToMember =
4359 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
4360
4361 VirtualPodLeafMap leafValues;
4362 for (const auto &[id, newMember] : idToMember) {
4363 leafValues[id] = preserveDiscardableAttrs(
4364 op, rewriter.create<MemberReadOp>(
4365 op.getLoc(), newMember.second, adaptor.getComponent(), newMember.first
4366 )
4367 );
4368 }
4369
4370 PodType podTy = llvm::cast<PodType>(op.getType());
4371 Value virtualPod = createVirtualPodPlaceholder(rewriter, op.getLoc(), podTy, leafValues);
4372 resolver.virtualPods[virtualPod] = std::move(leafValues);
4373 rewriter.replaceOp(op, virtualPod);
4374 }
4375};
4376
4381static bool tryCollectMaterializedSplitPodArrayLeafValues(
4382 Value arrayValue, ArrayType arrTy, ArrayRef<Type> splitTypes, SmallVectorImpl<Value> &leafArrays
4383) {
4384 auto cast = arrayValue.getDefiningOp<UnrealizedConversionCastOp>();
4385 size_t expectedOperands = splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0);
4386 if (!cast || cast->getNumResults() != 1 || cast.getResult(0).getType() != arrTy ||
4387 cast->getNumOperands() != expectedOperands) {
4388 return false;
4389 }
4390 if (needsPodArrayShapeCarrier(arrTy) &&
4391 cast.getOperand(splitTypes.size()).getType() != getPodArrayShapeCarrierType(arrTy)) {
4392 return false;
4393 }
4394
4396 cast.getOperands().take_front(splitTypes.size()), splitTypes, leafArrays
4397 );
4398}
4399
4405static bool tryCollectReadPodSplitPodArrayLeafValues(
4406 ReadPodOp readOp, ArrayType arrTy, ArrayRef<RecordChain> splitIds, ArrayRef<Type> splitTypes,
4407 const VirtualPodValueMap &virtualPods, SmallVectorImpl<Value> &leafArrays
4408) {
4409 auto tryCollectFromVirtualRead = [&](ReadPodOp sourceRead) {
4410 if (hasEarlierWrite(sourceRead)) {
4411 return false;
4412 }
4413
4414 const VirtualPodLeafMap *podLeafValues =
4415 lookupVirtualPodLeafMap(sourceRead.getPodRef(), virtualPods);
4416 if (!podLeafValues) {
4417 return false;
4418 }
4419
4420 SmallVector<Value> stagedLeafArrays;
4421 stagedLeafArrays.reserve(splitIds.size());
4422 for (const RecordChain &id : splitIds) {
4423 auto it = podLeafValues->find(id.withPrefix({sourceRead.getRecordNameAttr()}));
4424 if (it == podLeafValues->end() ||
4425 !typesUnify(it->second.getType(), getFlattenedTypeAlongPath(arrTy, id))) {
4426 return false;
4427 }
4428 stagedLeafArrays.push_back(it->second);
4429 }
4430 llvm::append_range(leafArrays, stagedLeafArrays);
4431 return true;
4432 };
4433
4434 if (WritePodOp writeOp = findNearestForwardableWrite(readOp)) {
4435 if (tryCollectMaterializedSplitPodArrayLeafValues(
4436 writeOp.getValue(), arrTy, splitTypes, leafArrays
4437 )) {
4438 return true;
4439 }
4440
4441 if (ReadPodOp writtenRead = peelUnifiableCasts(writeOp.getValue()).getDefiningOp<ReadPodOp>()) {
4442 if (tryCollectFromVirtualRead(writtenRead)) {
4443 return true;
4444 }
4445 }
4446 }
4447
4448 if (tryCollectFromVirtualRead(readOp)) {
4449 return true;
4450 }
4451
4452 return false;
4453}
4454
4456static bool resolveReadPodSplitPodArrayLeafValues(
4457 ReadPodOp readOp, ArrayType arrTy, ArrayRef<RecordChain> splitIds, ArrayRef<Type> splitTypes,
4458 const VirtualPodValueMap &virtualPods, DeferredPodArrayBackingMap &deferredPodArrays,
4459 Location loc, OpBuilder &bldr, SmallVectorImpl<Value> &leafArrays
4460) {
4461 if (tryCollectReadPodSplitPodArrayLeafValues(
4462 readOp, arrTy, splitIds, splitTypes, virtualPods, leafArrays
4463 )) {
4464 return true;
4465 }
4466
4467 if (!isFreshUnwrittenPodRead(readOp)) {
4468 return false;
4469 }
4470
4471 // Reuse one synthetic split-array backing per deferred field read so repeated users of the same
4472 // aggregate value continue to observe the same unwritten leaf storage and shared shape witness.
4473 DeferredPodArrayBacking &backing = materializeDeferredPodArrayBacking(
4474 readOp, arrTy, splitTypes, deferredPodArrays, loc, bldr,
4475 /*requireLeafArrays=*/true, /*requireShapeCarrier=*/true
4476 );
4477 leafArrays.assign(backing.leafArrays.begin(), backing.leafArrays.end());
4478
4479 return true;
4480}
4481
4483static Value tryResolveReadPodArrayShapeSource(
4484 ReadPodOp readOp, ArrayType arrTy, const VirtualPodValueMap &virtualPods, Location loc,
4485 RewriterBase &rewriter
4486) {
4487 auto tryGetVirtualShapeSource = [&](ReadPodOp sourceRead) -> Value {
4488 if (hasEarlierWrite(sourceRead)) {
4489 return {};
4490 }
4491
4492 const VirtualPodLeafMap *podLeafValues =
4493 lookupVirtualPodLeafMap(sourceRead.getPodRef(), virtualPods);
4494 if (!podLeafValues) {
4495 return {};
4496 }
4497
4498 SmallVector<StringAttr> carrierPath {
4499 sourceRead.getRecordNameAttr(), getPodArrayShapeCarrierMarker(rewriter.getContext())
4500 };
4501 if (auto carrierIt = podLeafValues->find(RecordChain(carrierPath, true));
4502 carrierIt != podLeafValues->end()) {
4503 return castValueToTypeIfNeeded(
4504 rewriter, loc, carrierIt->second, getPodArrayShapeCarrierType(arrTy)
4505 );
4506 }
4507
4508 size_t originalRank = arrTy.getDimensionSizes().size();
4509 SmallVector<RecordChain> splitIds;
4510 SmallVector<Type> splitTypes;
4511 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4512 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
4513 auto splitArrTy = llvm::dyn_cast<ArrayType>(splitType);
4514 if (!splitArrTy || splitArrTy.getDimensionSizes().size() != originalRank) {
4515 continue;
4516 }
4517
4518 auto valueIt = podLeafValues->find(id.withPrefix({sourceRead.getRecordNameAttr()}));
4519 if (valueIt != podLeafValues->end()) {
4520 return castValueToTypeIfNeeded(rewriter, loc, valueIt->second, splitArrTy);
4521 }
4522 }
4523
4524 return {};
4525 };
4526
4527 if (!hasEarlierWrite(readOp)) {
4528 if (Value shapeSource = tryGetVirtualShapeSource(readOp)) {
4529 return shapeSource;
4530 }
4531 }
4532
4533 if (WritePodOp writeOp = findNearestForwardableWrite(readOp)) {
4534 if (Value shapeSource = tryCollectDirectConvertedPodArrayShapeSource(
4535 writeOp.getValue(), arrTy, loc, rewriter
4536 )) {
4537 return shapeSource;
4538 }
4539 if (ReadPodOp writtenRead = peelUnifiableCasts(writeOp.getValue()).getDefiningOp<ReadPodOp>()) {
4540 if (Value shapeSource = tryGetVirtualShapeSource(writtenRead)) {
4541 return shapeSource;
4542 }
4543 }
4544 }
4545
4546 return {};
4547}
4548
4550static void eraseDeadDeferredFieldReadChain(ReadPodOp readOp, PatternRewriter &rewriter) {
4551 if (!readOp.getResult().use_empty()) {
4552 return;
4553 }
4554
4555 Value podRef = readOp.getPodRef();
4556 rewriter.eraseOp(readOp);
4557 if (podRef.use_empty()) {
4558 if (auto cast = podRef.getDefiningOp<UnrealizedConversionCastOp>()) {
4559 if (cast->getNumResults() == 1 && cast.getResult(0) == podRef) {
4560 rewriter.eraseOp(cast);
4561 }
4562 }
4563 }
4564}
4565
4567static bool getDeferredSplitPodArrayCastInfo(
4568 UnrealizedConversionCastOp op, ArrayType &arrTy, SmallVector<RecordChain> &splitIds,
4569 SmallVectorImpl<Type> &splitTypes
4570) {
4571 if (op->getNumOperands() != 1) {
4572 return false;
4573 }
4574
4575 arrTy = splittablePodArray(op.getOperand(0).getType());
4576 if (!arrTy) {
4577 return false;
4578 }
4579
4580 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4581 size_t expectedResults = splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0);
4582 if (op->getNumResults() != expectedResults) {
4583 return false;
4584 }
4585
4586 for (auto [result, splitType] :
4587 llvm::zip_equal(op.getResults().take_front(splitTypes.size()), splitTypes)) {
4588 if (result.getType() != splitType) {
4589 return false;
4590 }
4591 }
4592 return !needsPodArrayShapeCarrier(arrTy) ||
4593 op.getResult(splitTypes.size()).getType() == getPodArrayShapeCarrierType(arrTy);
4594}
4595
4597static LogicalResult resolveDeferredSplitPodArrayCast(
4598 UnrealizedConversionCastOp op, PatternRewriter &rewriter, Step3Resolver &resolver
4599) {
4600 ArrayType arrTy;
4601 SmallVector<RecordChain> splitIds;
4602 SmallVector<Type> splitTypes;
4603 if (!getDeferredSplitPodArrayCastInfo(op, arrTy, splitIds, splitTypes)) {
4604 return failure();
4605 }
4606
4607 ReadPodOp fieldRead = peelUnifiableCasts(op.getOperand(0)).getDefiningOp<ReadPodOp>();
4608 if (!fieldRead) {
4609 return failure();
4610 }
4611
4612 SmallVector<Value> splitLeafArrays;
4613 if (!resolveReadPodSplitPodArrayLeafValues(
4614 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, resolver.deferredPodArrays,
4615 op.getLoc(), rewriter, splitLeafArrays
4616 )) {
4617 return failure();
4618 }
4619
4620 SmallVector<Value> replacements(splitLeafArrays.begin(), splitLeafArrays.end());
4621 if (needsPodArrayShapeCarrier(arrTy)) {
4622 Value carrier = tryResolveReadPodArrayShapeSource(
4623 fieldRead, arrTy, resolver.virtualPods, op.getLoc(), rewriter
4624 );
4625 if (!carrier) {
4626 if (auto it = resolver.deferredPodArrays.find(fieldRead.getResult());
4627 it != resolver.deferredPodArrays.end() && it->second.shapeCarrier) {
4628 carrier = castValueToTypeIfNeeded(
4629 rewriter, op.getLoc(), it->second.shapeCarrier, getPodArrayShapeCarrierType(arrTy)
4630 );
4631 } else {
4632 carrier =
4633 materializeArrayLengthCarrier(fieldRead.getResult(), arrTy, op.getLoc(), rewriter);
4634 }
4635 }
4636 replacements.push_back(carrier);
4637 }
4638
4639 rewriter.replaceOp(op, replacements);
4640 eraseDeadDeferredFieldReadChain(fieldRead, rewriter);
4641 return success();
4642}
4643
4649class ResolvePodReadBackedArrayReadOp : public OpConversionPattern<ReadArrayOp> {
4650 Step3Resolver &resolver;
4651
4652public:
4653 ResolvePodReadBackedArrayReadOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4654 : OpConversionPattern<ReadArrayOp>(ctx), resolver(step3Resolver) {}
4655
4656 static bool canResolve(ReadArrayOp op, const Step3Resolver &resolver) {
4657 if (!shouldDeferPodArrayReadToStep3(op)) {
4658 return false;
4659 }
4660
4661 ArrayType arrTy = op.getArrRefType();
4662 auto fieldRead = llvm::cast<ReadPodOp>(op.getArrRef().getDefiningOp());
4663 SmallVector<RecordChain> splitIds;
4664 SmallVector<Type> splitTypes;
4665 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4666
4667 SmallVector<Value> ignoredLeafArrays;
4668 return tryCollectReadPodSplitPodArrayLeafValues(
4669 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, ignoredLeafArrays
4670 ) ||
4671 isFreshUnwrittenPodRead(fieldRead);
4672 }
4673
4674 LogicalResult matchAndRewrite(
4675 ReadArrayOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4676 ) const override {
4677 auto fieldRead = op.getArrRef().getDefiningOp<ReadPodOp>();
4678 if (!fieldRead) {
4679 return failure();
4680 }
4681
4682 ArrayType arrTy = op.getArrRefType();
4683 PodType podTy = llvm::cast<PodType>(arrTy.getElementType());
4684 SmallVector<RecordChain> splitIds;
4685 SmallVector<Type> splitTypes;
4686 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4687
4688 SmallVector<Value> splitLeafArrays;
4689 if (!resolveReadPodSplitPodArrayLeafValues(
4690 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods,
4691 resolver.deferredPodArrays, op.getLoc(), rewriter, splitLeafArrays
4692 )) {
4693 return failure();
4694 }
4695
4696 SmallVector<Value> indices(adaptor.getIndices().begin(), adaptor.getIndices().end());
4697 VirtualPodLeafMap leafValues;
4698 for (auto [id, splitType, leafArray] : llvm::zip_equal(splitIds, splitTypes, splitLeafArrays)) {
4699 Value leafValue = ArrayAccessOpInterface::genRead(rewriter, op.getLoc(), leafArray, indices);
4700 preserveDiscardableAttrs(op, leafValue.getDefiningOp());
4701 leafValues[id] = tagRaggedNestedLeafValue(
4702 rewriter, op.getLoc(), leafValue, getRaggedNestedLeafAttrName(arrTy, splitType)
4703 );
4704 }
4705
4706 Value virtualPod = createVirtualPodPlaceholder(rewriter, op.getLoc(), podTy, leafValues);
4707 resolver.virtualPods[virtualPod] = std::move(leafValues);
4708 rewriter.replaceOp(op, virtualPod);
4709 eraseDeadDeferredFieldReadChain(fieldRead, rewriter);
4710 return success();
4711 }
4712};
4713
4721class ResolveDeferredSplitPodArrayCastOp : public OpConversionPattern<UnrealizedConversionCastOp> {
4722 Step3Resolver &resolver;
4723
4724public:
4725 ResolveDeferredSplitPodArrayCastOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4726 : OpConversionPattern<UnrealizedConversionCastOp>(ctx), resolver(step3Resolver) {}
4727
4728 static bool canResolve(UnrealizedConversionCastOp op, const Step3Resolver &resolver) {
4729 ArrayType arrTy;
4730 SmallVector<RecordChain> splitIds;
4731 SmallVector<Type> splitTypes;
4732 if (!getDeferredSplitPodArrayCastInfo(op, arrTy, splitIds, splitTypes)) {
4733 return false;
4734 }
4735
4736 ReadPodOp fieldRead = peelUnifiableCasts(op.getOperand(0)).getDefiningOp<ReadPodOp>();
4737 if (!fieldRead) {
4738 return false;
4739 }
4740
4741 SmallVector<Value> ignoredLeafArrays;
4742 return tryCollectReadPodSplitPodArrayLeafValues(
4743 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, ignoredLeafArrays
4744 ) ||
4745 isFreshUnwrittenPodRead(fieldRead);
4746 }
4747
4748 LogicalResult matchAndRewrite(
4749 UnrealizedConversionCastOp op, OpAdaptor, ConversionPatternRewriter &rewriter
4750 ) const override {
4751 return resolveDeferredSplitPodArrayCast(op, rewriter, resolver);
4752 }
4753};
4754
4756class ResolveDeferredSplitPodArrayCastPrepass final
4757 : public OpRewritePattern<UnrealizedConversionCastOp> {
4758 Step3Resolver &resolver;
4759
4760public:
4761 ResolveDeferredSplitPodArrayCastPrepass(MLIRContext *ctx, Step3Resolver &step3Resolver)
4762 : OpRewritePattern<UnrealizedConversionCastOp>(ctx), resolver(step3Resolver) {}
4763
4764 LogicalResult
4765 matchAndRewrite(UnrealizedConversionCastOp op, PatternRewriter &rewriter) const override {
4766 return resolveDeferredSplitPodArrayCast(op, rewriter, resolver);
4767 }
4768};
4769
4771static LogicalResult splitVirtualPodEmitEquality(
4772 constrain::EmitEqualityOp op, RewriterBase &rewriter, Step3Resolver &resolver
4773) {
4774 return splitWholePodEmitEquality(
4775 op, rewriter, resolver.materializedLeaves, op.getOperation(), &resolver.virtualPods
4776 );
4777}
4778
4785static LogicalResult splitEarlierVirtualPodEqualitiesBeforeWriteInBlock(
4786 WritePodOp writeOp, RewriterBase &rewriter, Step3Resolver &resolver
4787) {
4788 Value writtenPod = peelVirtualPodCompatibilityCasts(writeOp.getPodRef());
4789
4790 SmallVector<constrain::EmitEqualityOp> equalityOps;
4791 for (Operation &candidate : llvm::make_early_inc_range(*writeOp->getBlock())) {
4792 if (&candidate == writeOp) {
4793 break;
4794 }
4795 candidate.walk<WalkOrder::PreOrder>([&equalityOps](constrain::EmitEqualityOp op) {
4796 if (getWholePodEqualityType(op)) {
4797 equalityOps.push_back(op);
4798 }
4799 });
4800 }
4801
4802 for (constrain::EmitEqualityOp equalityOp : equalityOps) {
4803 Value lhs = peelVirtualPodCompatibilityCasts(equalityOp.getLhs());
4804 Value rhs = peelVirtualPodCompatibilityCasts(equalityOp.getRhs());
4805 if (lhs != writtenPod && rhs != writtenPod) {
4806 continue;
4807 }
4808
4809 if (!lookupVirtualPodLeafMap(equalityOp.getLhs(), resolver.virtualPods) &&
4810 !lookupVirtualPodLeafMap(equalityOp.getRhs(), resolver.virtualPods)) {
4811 continue;
4812 }
4813
4814 rewriter.setInsertionPoint(equalityOp);
4815 if (failed(splitVirtualPodEmitEquality(equalityOp, rewriter, resolver))) {
4816 return failure();
4817 }
4818 }
4819 return success();
4820}
4821
4826class ResolveVirtualPodWriteOp : public OpConversionPattern<WritePodOp> {
4827 Step3Resolver &resolver;
4828
4829public:
4830 ResolveVirtualPodWriteOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4831 : OpConversionPattern<WritePodOp>(ctx), resolver(step3Resolver) {}
4832
4833 LogicalResult matchAndRewrite(
4834 WritePodOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4835 ) const override {
4836 auto it = lookupVirtualPodLeafMapIt(op.getPodRef(), resolver.virtualPods);
4837 if (it == resolver.virtualPods.end() || isInsideSupportedScfRegion(op.getOperation())) {
4838 return failure();
4839 }
4840
4841 if (failed(splitEarlierVirtualPodEqualitiesBeforeWriteInBlock(op, rewriter, resolver))) {
4842 return failure();
4843 }
4844
4845 Type recordType =
4846 llvm::cast<PodType>(op.getPodRefType()).getRecordMap().lookup(op.getRecordName());
4847 assert(recordType && "record must exist in POD type");
4848 updateVirtualPodRecordLeafValues(
4849 op.getLoc(), op.getRecordNameAttr(), recordType, adaptor.getValue(), resolver.virtualPods,
4850 rewriter, it->second
4851 );
4852 rewriter.eraseOp(op);
4853 return success();
4854 }
4855};
4856
4861class ResolveVirtualPodReadOp : public OpConversionPattern<ReadPodOp> {
4862 Step3Resolver &resolver;
4863
4864public:
4865 ResolveVirtualPodReadOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4866 : OpConversionPattern<ReadPodOp>(ctx), resolver(step3Resolver) {}
4867
4868 LogicalResult
4869 matchAndRewrite(ReadPodOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
4870 if (hasEarlierWrite(op) || findNearestForwardableWrite(op)) {
4871 return failure();
4872 }
4873
4874 const VirtualPodLeafMap *leafValues =
4875 lookupVirtualPodLeafMap(op.getPodRef(), resolver.virtualPods);
4876 if (!leafValues) {
4877 return failure();
4878 }
4879
4880 SmallVector<StringAttr> prefix {op.getRecordNameAttr()};
4881 Type recordType =
4882 llvm::cast<PodType>(op.getPodRefType()).getRecordMap().lookup(op.getRecordName());
4883 assert(recordType && "record must exist in POD type");
4884
4885 if (PodType nestedPodTy = llvm::dyn_cast<PodType>(recordType)) {
4886 VirtualPodLeafMap nestedLeafValues;
4887 SmallVector<StringAttr> nestedRecordChain;
4888 forEachPodLeaf(nestedPodTy, nestedRecordChain, [&](const RecordChain &id, Type) {
4889 nestedLeafValues[id] = leafValues->at(id.withPrefix(prefix));
4890 });
4891 Value virtualPod =
4892 createVirtualPodPlaceholder(rewriter, op.getLoc(), nestedPodTy, nestedLeafValues);
4893 resolver.virtualPods[virtualPod] = std::move(nestedLeafValues);
4894 rewriter.replaceOp(op, virtualPod);
4895 return success();
4896 }
4897
4898 if (splittablePodArray(recordType)) {
4899 return failure();
4900 }
4901
4902 rewriter.replaceOp(
4903 op, castValueToTypeIfNeeded(
4904 rewriter, op.getLoc(), leafValues->at(RecordChain(prefix)), recordType
4905 )
4906 );
4907 return success();
4908 }
4909};
4910
4912static void rehydrateVirtualPodPlaceholders(ModuleOp modOp, VirtualPodValueMap &virtualPods) {
4913 modOp.walk([&virtualPods](UnrealizedConversionCastOp castOp) {
4914 if (castOp->getNumResults() != 1) {
4915 return;
4916 }
4917
4918 PodType podTy = llvm::dyn_cast<PodType>(castOp.getResult(0).getType());
4919 if (!podTy) {
4920 return;
4921 }
4922
4923 VirtualPodLeafMap leafValues;
4924 SmallVector<StringAttr> recordChain;
4925 auto operandIt = castOp.getOperands().begin();
4926 bool matchesVirtualPlaceholder = true;
4927 forEachPodLeaf(podTy, recordChain, [&](const RecordChain &id, Type leafType) {
4928 if (!matchesVirtualPlaceholder || operandIt == castOp.getOperands().end() ||
4929 !typesUnify((*operandIt).getType(), leafType)) {
4930 matchesVirtualPlaceholder = false;
4931 return;
4932 }
4933 leafValues[id] = *operandIt++;
4934 });
4935
4936 if (matchesVirtualPlaceholder && operandIt == castOp.getOperands().end()) {
4937 virtualPods[castOp.getResult(0)] = std::move(leafValues);
4938 }
4939 });
4940}
4941
4943class SplitVirtualPodInEmitEqualityPattern final
4944 : public OpRewritePattern<constrain::EmitEqualityOp> {
4945 Step3Resolver &resolver;
4946
4947public:
4948 SplitVirtualPodInEmitEqualityPattern(MLIRContext *ctx, Step3Resolver &step3Resolver)
4949 : OpRewritePattern<constrain::EmitEqualityOp>(ctx), resolver(step3Resolver) {}
4950
4951 LogicalResult
4952 matchAndRewrite(constrain::EmitEqualityOp op, PatternRewriter &rewriter) const override {
4953 if (!getWholePodEqualityType(op)) {
4954 return failure();
4955 }
4956
4957 return splitVirtualPodEmitEquality(op, rewriter, resolver);
4958 }
4959};
4960
4968struct PromotedFunctionArgCast {
4969 unsigned argIndex;
4970 SmallVector<Type> resultTypes;
4971 SmallVector<UnrealizedConversionCastOp> casts;
4972};
4973
4979struct PromotedFunctionSignature {
4980 FuncDefOp func;
4981 Operation *funcOp;
4982 unsigned oldInputCount;
4983 SmallVector<PromotedFunctionArgCast> argCasts;
4984};
4985
4986static const PromotedFunctionArgCast *
4987findPromotedArgCast(ArrayRef<PromotedFunctionArgCast> argCasts, unsigned argIndex) {
4988 auto it = llvm::find_if(argCasts, [argIndex](const PromotedFunctionArgCast &argCast) {
4989 return argCast.argIndex == argIndex;
4990 });
4991 return it == argCasts.end() ? nullptr : &*it;
4992}
4993
5000static std::optional<PromotedFunctionArgCast> getPromotableFunctionArgCast(BlockArgument arg) {
5001 PromotedFunctionArgCast promoted;
5002 promoted.argIndex = arg.getArgNumber();
5003 bool initialized = false;
5004
5005 for (OpOperand &use : arg.getUses()) {
5006 auto castOp = dyn_cast<UnrealizedConversionCastOp>(use.getOwner());
5007 if (!castOp || castOp->getNumOperands() != 1 || castOp.getOperand(0) != arg ||
5008 castOp->getNumResults() <= 1) {
5009 return std::nullopt;
5010 }
5011
5012 SmallVector<Type> castResultTypes(castOp.getResultTypes());
5013 if (!initialized) {
5014 promoted.resultTypes = std::move(castResultTypes);
5015 initialized = true;
5016 } else if (promoted.resultTypes != castResultTypes) {
5017 return std::nullopt;
5018 }
5019 promoted.casts.push_back(castOp);
5020 }
5021
5022 return initialized ? std::optional<PromotedFunctionArgCast>(std::move(promoted)) : std::nullopt;
5023}
5024
5030static ArrayAttr expandArgAttrsForPromotedFunctionArgCasts(
5031 FuncDefOp func, ArrayRef<PromotedFunctionArgCast> argCasts
5032) {
5033 ArrayAttr origAttrs = func.getArgAttrsAttr();
5034 if (!origAttrs) {
5035 return nullptr;
5036 }
5037
5038 llvm::StringSet<> usedNames;
5039 for (auto [i, attr] : llvm::enumerate(origAttrs)) {
5040 if (findPromotedArgCast(argCasts, i)) {
5041 continue;
5042 }
5043 auto dictAttr = dyn_cast<DictionaryAttr>(attr);
5044 if (!dictAttr) {
5045 continue;
5046 }
5047 if (auto nameAttr = dyn_cast_if_present<StringAttr>(dictAttr.get(ARG_NAME_ATTR_NAME))) {
5048 usedNames.insert(nameAttr.getValue());
5049 }
5050 }
5051
5052 SmallVector<Attribute> newAttrs;
5053 for (auto [i, attr] : llvm::enumerate(origAttrs)) {
5054 const PromotedFunctionArgCast *argCast = findPromotedArgCast(argCasts, i);
5055 if (!argCast) {
5056 newAttrs.push_back(attr);
5057 continue;
5058 }
5059
5060 auto dictAttr = mlir::cast<DictionaryAttr>(attr);
5061 auto nameAttr = dyn_cast_if_present<StringAttr>(dictAttr.get(ARG_NAME_ATTR_NAME));
5062 if (!nameAttr) {
5063 newAttrs.append(argCast->resultTypes.size(), attr);
5064 continue;
5065 }
5066
5067 llvm::StringRef baseName = nameAttr.getValue();
5068 for (unsigned splitIdx = 0, e = argCast->resultTypes.size(); splitIdx < e; ++splitIdx) {
5069 std::string desiredName =
5070 splitIdx == 0 ? baseName.str() : (baseName + "#" + llvm::Twine(splitIdx)).str();
5071 newAttrs.push_back(
5072 withFunctionArgNameAttr(dictAttr, reserveUniqueAttrName(usedNames, desiredName))
5073 );
5074 }
5075 }
5076 return ArrayAttr::get(func.getContext(), newAttrs);
5077}
5078
5087static std::optional<PromotedFunctionSignature> promoteFunctionArgCasts(FuncDefOp func) {
5088 if (func.isExternal()) {
5089 return std::nullopt;
5090 }
5091
5092 Block &entryBlock = func.getBody().front();
5093 SmallVector<PromotedFunctionArgCast> argCasts;
5094 for (BlockArgument arg : entryBlock.getArguments()) {
5095 if (auto promoted = getPromotableFunctionArgCast(arg)) {
5096 argCasts.push_back(std::move(*promoted));
5097 }
5098 }
5099 if (argCasts.empty()) {
5100 return std::nullopt;
5101 }
5102
5103 FunctionType oldFuncTy = func.getFunctionType();
5104 SmallVector<Type> newInputs;
5105 for (auto [i, inputType] : llvm::enumerate(oldFuncTy.getInputs())) {
5106 if (const PromotedFunctionArgCast *argCast = findPromotedArgCast(argCasts, i)) {
5107 llvm::append_range(newInputs, argCast->resultTypes);
5108 } else {
5109 newInputs.push_back(inputType);
5110 }
5111 }
5112
5113 ArrayAttr newArgAttrs = expandArgAttrsForPromotedFunctionArgCasts(func, argCasts);
5114 func.setFunctionType(FunctionType::get(func.getContext(), newInputs, oldFuncTy.getResults()));
5115 if (newArgAttrs) {
5116 func.setArgAttrsAttr(newArgAttrs);
5117 }
5118
5119 for (const PromotedFunctionArgCast &argCast : llvm::reverse(argCasts)) {
5120 BlockArgument oldArg = entryBlock.getArgument(argCast.argIndex);
5121 SmallVector<BlockArgument> newArgs;
5122 newArgs.reserve(argCast.resultTypes.size());
5123 unsigned nextArgIdx = argCast.argIndex + 1;
5124 for (Type resultType : argCast.resultTypes) {
5125 newArgs.push_back(entryBlock.insertArgument(nextArgIdx, resultType, oldArg.getLoc()));
5126 ++nextArgIdx;
5127 }
5128
5129 for (UnrealizedConversionCastOp castOp : argCast.casts) {
5130 for (auto [result, newArg] : llvm::zip_equal(castOp.getResults(), newArgs)) {
5131 result.replaceAllUsesWith(newArg);
5132 }
5133 castOp.erase();
5134 }
5135 entryBlock.eraseArgument(argCast.argIndex);
5136 }
5137
5138 return PromotedFunctionSignature {
5139 .func = func,
5140 .funcOp = func.getOperation(),
5141 .oldInputCount = oldFuncTy.getNumInputs(),
5142 .argCasts = std::move(argCasts)
5143 };
5144}
5145
5152static LogicalResult updateCallsForPromotedFunctionArgCasts(
5153 ModuleOp modOp, SymbolTableCollection &symTables,
5154 ArrayRef<PromotedFunctionSignature> promotedSignatures
5155) {
5156 if (promotedSignatures.empty()) {
5157 return success();
5158 }
5159
5160 DenseMap<Operation *, const PromotedFunctionSignature *> promotedByFunc;
5161 for (const PromotedFunctionSignature &signature : promotedSignatures) {
5162 promotedByFunc[signature.funcOp] = &signature;
5163 }
5164
5165 OpBuilder builder(modOp.getContext());
5166 SmallVector<CallOp> calls;
5167 modOp.walk([&calls](CallOp callOp) { calls.push_back(callOp); });
5168 for (CallOp callOp : calls) {
5169 FailureOr<SymbolLookupResult<FuncDefOp>> targetRes = callOp.getCalleeTarget(symTables);
5170 if (failed(targetRes)) {
5171 return failure();
5172 }
5173
5174 auto signatureIt = promotedByFunc.find(targetRes->get().getOperation());
5175 if (signatureIt == promotedByFunc.end()) {
5176 continue;
5177 }
5178
5179 const PromotedFunctionSignature &signature = *signatureIt->second;
5180 if (callOp.getArgOperands().size() != signature.oldInputCount) {
5181 return callOp.emitOpError("argument count does not match pre-promotion callee signature");
5182 }
5183
5184 SmallVector<Value> newOperands;
5185 builder.setInsertionPoint(callOp);
5186 for (auto [i, operand] : llvm::enumerate(callOp.getArgOperands())) {
5187 if (const PromotedFunctionArgCast *argCast = findPromotedArgCast(signature.argCasts, i)) {
5188 auto castOp = builder.create<UnrealizedConversionCastOp>(
5189 callOp.getLoc(), TypeRange(argCast->resultTypes), operand
5190 );
5191 llvm::append_range(newOperands, castOp.getResults());
5192 } else {
5193 newOperands.push_back(operand);
5194 }
5195 }
5196 callOp.getArgOperandsMutable().assign(ValueRange(newOperands));
5197 }
5198 return success();
5199}
5200
5206static LogicalResult
5207promoteFunctionArgCastsToSignature(ModuleOp modOp, SymbolTableCollection &symTables) {
5208 SmallVector<FuncDefOp> funcs;
5209 modOp.walk([&funcs](FuncDefOp func) { funcs.push_back(func); });
5210
5211 for (unsigned promotionRounds = 1;; ++promotionRounds) {
5212 SmallVector<PromotedFunctionSignature> promotedSignatures;
5213 for (FuncDefOp func : funcs) {
5214 if (auto promoted = promoteFunctionArgCasts(func)) {
5215 promotedSignatures.push_back(std::move(*promoted));
5216 }
5217 }
5218 if (promotedSignatures.empty()) {
5219 return success();
5220 }
5221 if (failed(updateCallsForPromotedFunctionArgCasts(modOp, symTables, promotedSignatures))) {
5222 return failure();
5223 }
5224 if (promotionRounds % 64 == 0) {
5225 llvm::outs() << "function argument cast promotion has run " << promotionRounds
5226 << " rounds without reaching a fixpoint; continuing...\n";
5227 }
5228 }
5229}
5230
5231void Step3Resolver::rehydrateVirtualPodPlaceholders(ModuleOp modOp) {
5232 ::rehydrateVirtualPodPlaceholders(modOp, virtualPods);
5233}
5234
5235void Step3Resolver::addPreConversionPatterns(RewritePatternSet &patterns) {
5236 patterns.add<ResolveDeferredSplitPodArrayCastPrepass>(patterns.getContext(), *this);
5237}
5238
5239void Step3Resolver::addConversionPatterns(
5240 RewritePatternSet &patterns, SymbolTableCollection &symTables,
5241 const MemberReplacementMap &memberRepMap
5242) {
5243 patterns.add<SplitInitFromNewPodOp>(patterns.getContext());
5244 patterns.add<SplitPodElementCreateArrayOp>(patterns.getContext(), *this);
5245 patterns.add<SplitPodInFuncDefOp, SplitPodInReturnOp, SplitPodInCallOp>(
5246 patterns.getContext(), *this
5247 );
5248 patterns.add<SplitPodInMemberWriteOp, SplitPodInMemberReadOp>(
5249 patterns.getContext(), symTables, memberRepMap, *this
5250 );
5251 patterns.add<RejectRaggedNestedLeafArrayLengthOp>(patterns.getContext());
5252 patterns.add<ResolvePodReadBackedArrayReadOp>(patterns.getContext(), *this);
5253 patterns.add<ResolvePodReadBackedArrayLengthOp>(patterns.getContext(), *this);
5254 patterns.add<ResolveDeferredSplitPodArrayCastOp>(patterns.getContext(), *this);
5255 patterns.add<ResolveVirtualPodWriteOp>(patterns.getContext(), *this);
5256 patterns.add<ResolveVirtualPodReadOp>(patterns.getContext(), *this);
5257}
5258
5259void Step3Resolver::addLateResolutionPatterns(RewritePatternSet &patterns) {
5260 patterns.add<ResolvePodReadBackedArrayReadOp>(patterns.getContext(), *this);
5261 patterns.add<ResolvePodReadBackedArrayLengthOp>(patterns.getContext(), *this);
5262 patterns.add<ResolveDeferredSplitPodArrayCastOp>(patterns.getContext(), *this);
5263 patterns.add<ResolveVirtualPodWriteOp>(patterns.getContext(), *this);
5264 patterns.add<ResolveVirtualPodReadOp>(patterns.getContext(), *this);
5265}
5266
5267void Step3Resolver::addPostConversionPatterns(RewritePatternSet &patterns) {
5268 patterns.add<SplitVirtualPodInEmitEqualityPattern>(patterns.getContext(), *this);
5269}
5270
5271void Step3Resolver::configureLateVirtualPodLegality(ConversionTarget &target) const {
5272 target.addDynamicallyLegalOp<WritePodOp>([this](WritePodOp op) {
5273 return !lookupVirtualPodLeafMap(op.getPodRef(), virtualPods) ||
5274 isInsideSupportedScfRegion(op.getOperation());
5275 });
5276 target.addDynamicallyLegalOp<ArrayLengthOp>([](ArrayLengthOp op) {
5277 return RejectRaggedNestedLeafArrayLengthOp::legal(op) &&
5278 ResolvePodReadBackedArrayLengthOp::legal(op);
5279 });
5280 target.addDynamicallyLegalOp<ReadArrayOp>([this](ReadArrayOp op) {
5281 return !ResolvePodReadBackedArrayReadOp::canResolve(op, *this);
5282 });
5283 target.addDynamicallyLegalOp<UnrealizedConversionCastOp>([this](UnrealizedConversionCastOp op) {
5284 return !ResolveDeferredSplitPodArrayCastOp::canResolve(op, *this);
5285 });
5286 target.addDynamicallyLegalOp<ReadPodOp>([this](ReadPodOp op) {
5287 return !canResolveVirtualPodRead(op, virtualPods);
5288 });
5289}
5290
5291bool Step3Resolver::hasResolvableLateVirtualPodOps(ModuleOp modOp) const {
5292 return walkContainsMatch<Operation *>(*modOp, [this](Operation *op) {
5293 return TypeSwitch<Operation *, bool>(op)
5294 .Case<WritePodOp>([this](auto writeOp) {
5295 return lookupVirtualPodLeafMap(writeOp.getPodRef(), virtualPods) &&
5296 !isInsideSupportedScfRegion(writeOp.getOperation());
5297 })
5298 .Case<ReadPodOp>([this](auto readOp) {
5299 return canResolveVirtualPodRead(readOp, virtualPods);
5300 })
5301 .Case<ReadArrayOp>([this](auto readOp) {
5302 return ResolvePodReadBackedArrayReadOp::canResolve(readOp, *this);
5303 })
5304 .Case<ArrayLengthOp>([](auto lenOp) {
5305 return !ResolvePodReadBackedArrayLengthOp::legal(lenOp);
5306 })
5307 .Case<UnrealizedConversionCastOp>([this](auto castOp) {
5308 return ResolveDeferredSplitPodArrayCastOp::canResolve(castOp, *this);
5309 }).Default([](Operation *) { return false; });
5310 });
5311}
5312
5313void Step3Resolver::materializeRemainingVirtualPods(ModuleOp modOp) {
5314 OpBuilder builder(modOp.getContext());
5315 modOp.walk([this, &builder](NewPodOp newPod) {
5316 auto it = virtualPods.find(newPod.getResult());
5317 if (it == virtualPods.end() || newPod.use_empty()) {
5318 return;
5319 }
5320 builder.setInsertionPointAfter(findVirtualPodMaterializationAnchor(newPod, it->second));
5321 materializeVirtualPod(builder, newPod, it->second);
5322 });
5323}
5324
5327static LogicalResult
5328step3(ModuleOp modOp, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap) {
5329 MLIRContext *ctx = modOp.getContext();
5330 Step3Resolver resolver;
5331 resolver.rehydrateVirtualPodPlaceholders(modOp);
5332
5333 RewritePatternSet preConversionPatterns(ctx);
5334 resolver.addPreConversionPatterns(preConversionPatterns);
5335 if (failed(applyPatternsGreedily(
5336 modOp->getRegion(0), std::move(preConversionPatterns),
5337 GreedyRewriteConfig {.fold = false, .cseConstants = false}
5338 ))) {
5339 return failure();
5340 }
5341
5342 RewritePatternSet patterns(ctx);
5343 resolver.addConversionPatterns(patterns, symTables, memberRepMap);
5344
5345 ConversionTarget target(*ctx);
5346 baseTargetSetup(target);
5347 target.addDynamicallyLegalOp<NewPodOp>(SplitInitFromNewPodOp::legal);
5348 target.addDynamicallyLegalOp<CreateArrayOp>(SplitPodElementCreateArrayOp::legal);
5349 target.addDynamicallyLegalOp<FuncDefOp>(SplitPodInFuncDefOp::legal);
5350 target.addDynamicallyLegalOp<ReturnOp>(SplitPodInReturnOp::legal);
5351 target.addDynamicallyLegalOp<CallOp>(SplitPodInCallOp::legal);
5352 target.addDynamicallyLegalOp<MemberWriteOp>(SplitPodInMemberWriteOp::legal);
5353 target.addDynamicallyLegalOp<MemberReadOp>(SplitPodInMemberReadOp::legal);
5354 resolver.configureLateVirtualPodLegality(target);
5355
5356 LLVM_DEBUG(llvm::dbgs() << "Begin step 3: update/split other pod ops\n";);
5357 if (failed(applyFullConversion(modOp, target, std::move(patterns)))) {
5358 return failure();
5359 }
5360
5361 // Step 3's full conversion can expose a second wave of virtual-POD cleanup
5362 // opportunities. For example, resolving one deferred array bridge can make a
5363 // previously-illegal pod.read/pod.write or cast finally resolvable. Scan for
5364 // exactly those ops here so we can drive a small fixpoint loop before the
5365 // final materialization stage.
5366 // Rebuild the late-resolution target and pattern set each round because the
5367 // set of remaining virtual-POD placeholders shrinks as patterns fire. This
5368 // phase intentionally uses partial conversion: each round resolves whatever
5369 // is now legal to lower, then we rescan for any newly-exposed opportunities.
5370 auto runLateResolutionRound = [&]() {
5371 RewritePatternSet lateResolutionPatterns(ctx);
5372 resolver.addLateResolutionPatterns(lateResolutionPatterns);
5373
5374 ConversionTarget lateResolutionTarget(*ctx);
5375 baseTargetSetup(lateResolutionTarget);
5376 lateResolutionTarget.addLegalOp<ModuleOp>();
5377 resolver.configureLateVirtualPodLegality(lateResolutionTarget);
5378
5379 return applyPartialConversion(modOp, lateResolutionTarget, std::move(lateResolutionPatterns));
5380 };
5381
5382 // Iterate to a fixpoint before post-processing materializes any surviving
5383 // virtual PODs. A bounded loop guards against accidental non-progress.
5384 for (unsigned lateResolutionRounds = 0;
5385 lateResolutionRounds < 64 && resolver.hasResolvableLateVirtualPodOps(modOp);
5386 ++lateResolutionRounds) {
5387 if (failed(runLateResolutionRound())) {
5388 return failure();
5389 }
5390 }
5391 if (resolver.hasResolvableLateVirtualPodOps(modOp)) {
5392 modOp.emitError("late virtual POD resolution did not reach a fixpoint");
5393 return failure();
5394 }
5395
5396 RewritePatternSet postConversionPatterns(ctx);
5397 resolver.addPostConversionPatterns(postConversionPatterns);
5398 if (failed(applyPatternsGreedily(
5399 modOp->getRegion(0), std::move(postConversionPatterns),
5400 GreedyRewriteConfig {.fold = false, .cseConstants = false}
5401 ))) {
5402 return failure();
5403 }
5404 if (failed(promoteFunctionArgCastsToSignature(modOp, symTables))) {
5405 return failure();
5406 }
5407
5408 resolver.materializeRemainingVirtualPods(modOp);
5409
5410 bool erasedDeadPlaceholderOps = false;
5411 do {
5412 SmallVector<Operation *> deadPlaceholderOps;
5413 modOp->walk<WalkOrder::PostOrder>([&deadPlaceholderOps](Operation *op) {
5414 if (auto readOp = llvm::dyn_cast<ReadPodOp>(op)) {
5415 if (readOp.getResult().use_empty()) {
5416 deadPlaceholderOps.push_back(op);
5417 }
5418 return;
5419 }
5420
5421 if (auto castOp = llvm::dyn_cast<UnrealizedConversionCastOp>(op)) {
5422 if (llvm::all_of(castOp.getResults(), [](Value result) { return result.use_empty(); })) {
5423 deadPlaceholderOps.push_back(op);
5424 }
5425 }
5426 });
5427 for (Operation *op : deadPlaceholderOps) {
5428 op->erase();
5429 }
5430 erasedDeadPlaceholderOps = !deadPlaceholderOps.empty();
5431 } while (erasedDeadPlaceholderOps);
5432
5433 SmallVector<Operation *> deadOps;
5434 modOp->walk<WalkOrder::PostOrder>([&](Operation *op) {
5435 if (op != modOp.getOperation() && isOpTriviallyDead(op)) {
5436 deadOps.push_back(op);
5437 }
5438 });
5439 for (Operation *op : deadOps) {
5440 op->erase();
5441 }
5442 return success();
5443}
5444
5449static bool isValueDefinedInside(Operation *ancestor, Value value) {
5450 if (Operation *defOp = value.getDefiningOp()) {
5451 return ancestor->isAncestor(defOp);
5452 }
5453
5454 auto blockArg = llvm::dyn_cast<BlockArgument>(value);
5455 Operation *parentOp = blockArg.getOwner()->getParentOp();
5456 return parentOp && ancestor->isAncestor(parentOp);
5457}
5458
5460static WritePodOp findPrecedingWriteForIfRead(ReadPodOp readOp) {
5461 auto ifOp = readOp->getParentOfType<scf::IfOp>();
5462 if (!ifOp || readOp->getBlock()->getParentOp() != ifOp.getOperation()) {
5463 return nullptr;
5464 }
5465 if (hasEarlierWriteInBlock(readOp)) {
5466 return nullptr;
5467 }
5468
5469 Block *ifBlock = ifOp->getBlock();
5470 if (!ifBlock) {
5471 return nullptr;
5472 }
5473
5474 Value podRef = readOp.getPodRef();
5475 StringAttr recordName = readOp.getRecordNameAttr();
5476 WritePodOp replacement = nullptr;
5477 for (Operation &op : *ifBlock) {
5478 if (&op == ifOp.getOperation()) {
5479 break;
5480 }
5481
5482 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5483 if (isSamePodRecord(writeOp, podRef, recordName)) {
5484 replacement = writeOp;
5485 }
5486 continue;
5487 }
5488
5489 if (hasNestedWriteToRecord(op, podRef, recordName)) {
5490 replacement = nullptr;
5491 }
5492 }
5493
5494 return replacement;
5495}
5496
5498class FoldReadAfterWriteInBlockPattern final : public OpRewritePattern<ReadPodOp> {
5499public:
5500 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5501
5502 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5503 if (WritePodOp writeOp = findNearestForwardableWriteInBlock(readOp)) {
5504 rewriter.replaceOp(readOp, writeOp.getValue());
5505 return success();
5506 }
5507 return failure();
5508 }
5509};
5510
5512class ReplaceIfReadPattern final : public OpRewritePattern<ReadPodOp> {
5513public:
5514 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5515
5516 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5517 auto ifOp = readOp->getParentOfType<scf::IfOp>();
5518 if (!ifOp || readOp->getBlock()->getParentOp() != ifOp.getOperation()) {
5519 return failure();
5520 }
5521 if (isValueDefinedInside(ifOp, readOp.getPodRef()) || hasEarlierWriteInBlock(readOp)) {
5522 return failure();
5523 }
5524
5525 if (WritePodOp writeOp = findPrecedingWriteForIfRead(readOp)) {
5526 rewriter.replaceOp(readOp, writeOp.getValue());
5527 return success();
5528 }
5529
5530 rewriter.setInsertionPoint(ifOp);
5531 rewriter.replaceOp(
5532 readOp, genRead(rewriter, readOp.getLoc(), readOp.getPodRef(), readOp.getRecordNameAttr())
5533 .getResult()
5534 );
5535 return success();
5536 }
5537};
5538
5552class FoldIfCarriedPodReadAfterWritePattern final : public OpRewritePattern<ReadPodOp> {
5553public:
5554 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5555
5556 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5557 auto podRes = dyn_cast<OpResult>(readOp.getPodRef());
5558 if (!podRes) {
5559 return failure();
5560 }
5561
5562 auto ifOp = dyn_cast<scf::IfOp>(podRes.getOwner());
5563 if (!ifOp) {
5564 return failure();
5565 }
5566
5567 auto writeOp = dyn_cast_or_null<WritePodOp>(readOp->getPrevNode());
5568 if (!writeOp || writeOp.getRecordNameAttr() != readOp.getRecordNameAttr()) {
5569 return failure();
5570 }
5571
5572 auto valueRes = dyn_cast<OpResult>(writeOp.getValue());
5573 if (!valueRes || valueRes.getOwner() != ifOp.getOperation()) {
5574 return failure();
5575 }
5576
5577 Value carriedPod = writeOp.getPodRef();
5578 unsigned podResultIndex = podRes.getResultNumber();
5579
5580 auto thenYield = dyn_cast<scf::YieldOp>(ifOp.thenBlock()->getTerminator());
5581 if (!thenYield || thenYield.getOperand(podResultIndex) != carriedPod) {
5582 return failure();
5583 }
5584
5585 Region &elseRegion = ifOp.getElseRegion();
5586 if (Block *elseBlock = elseRegion.empty() ? nullptr : &elseRegion.front()) {
5587 auto elseYield = dyn_cast<scf::YieldOp>(elseBlock->getTerminator());
5588 if (!elseYield || elseYield.getOperand(podResultIndex) != carriedPod) {
5589 return failure();
5590 }
5591 }
5592
5593 rewriter.replaceOp(readOp, valueRes);
5594 return success();
5595 }
5596};
5597
5602struct IfWriteSlot {
5603 Value podRef;
5604 StringAttr recordName;
5605 Type type;
5606 WritePodOp thenWrite;
5607 WritePodOp elseWrite;
5608 Value incomingValue;
5609};
5610
5612static IfWriteSlot *
5613lookupSlot(SmallVectorImpl<IfWriteSlot> &slots, Value podRef, StringAttr recordName) {
5614 for (IfWriteSlot &slot : slots) {
5615 if (slot.podRef == podRef && slot.recordName == recordName) {
5616 return &slot;
5617 }
5618 }
5619 return nullptr;
5620}
5621
5623static IfWriteSlot &getOrCreateSlot(
5624 SmallVectorImpl<IfWriteSlot> &slots, Value podRef, StringAttr recordName, Type type
5625) {
5626 if (IfWriteSlot *slot = lookupSlot(slots, podRef, recordName)) {
5627 return *slot;
5628 }
5629 slots.push_back(IfWriteSlot {podRef, recordName, type, nullptr, nullptr, Value()});
5630 return slots.back();
5631}
5632
5634static Block *getElseBlockOrNull(scf::IfOp ifOp) {
5635 return ifOp.getElseRegion().empty() ? nullptr : &ifOp.getElseRegion().front();
5636}
5637
5639static void
5640collectDirectWrites(Block *block, bool isThenBlock, SmallVectorImpl<IfWriteSlot> &slots) {
5641 if (!block) {
5642 return;
5643 }
5644
5645 for (Operation &op : *block) {
5646 if (op.hasTrait<OpTrait::IsTerminator>()) {
5647 break;
5648 }
5649
5650 auto writeOp = dyn_cast<WritePodOp>(&op);
5651 if (!writeOp) {
5652 continue;
5653 }
5654
5655 IfWriteSlot &slot = getOrCreateSlot(
5656 slots, writeOp.getPodRef(), writeOp.getRecordNameAttr(), writeOp.getValue().getType()
5657 );
5658 if (isThenBlock) {
5659 slot.thenWrite = writeOp;
5660 } else {
5661 slot.elseWrite = writeOp;
5662 }
5663 }
5664}
5665
5670static bool branchSlotCanBeLifted(Block *block, Value podRef, StringAttr recordName) {
5671 if (!block) {
5672 return true;
5673 }
5674
5675 bool seenDirectWrite = false;
5676 for (Operation &op : *block) {
5677 if (op.hasTrait<OpTrait::IsTerminator>()) {
5678 return true;
5679 }
5680
5681 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5682 if (isSamePodRecord(writeOp, podRef, recordName)) {
5683 seenDirectWrite = true;
5684 continue;
5685 }
5686 }
5687
5688 if (hasNestedWriteToRecord(op, podRef, recordName)) {
5689 return false;
5690 }
5691 if (seenDirectWrite && (hasReadFromRecord(op, podRef, recordName) || hasValueUse(op, podRef))) {
5692 return false;
5693 }
5694 }
5695 return true;
5696}
5697
5699static bool isLiftedWrite(Operation &op, ArrayRef<IfWriteSlot> slots) {
5700 auto writeOp = dyn_cast<WritePodOp>(&op);
5701 return writeOp && llvm::any_of(slots, [&writeOp](const IfWriteSlot &slot) {
5702 return isSamePodRecord(writeOp, slot.podRef, slot.recordName);
5703 });
5704}
5705
5707static scf::YieldOp getYieldOp(Block &block) {
5708 auto yieldOp = dyn_cast<scf::YieldOp>(block.getTerminator());
5709 assert(yieldOp && "expected scf.if branch to terminate with scf.yield");
5710 return yieldOp;
5711}
5712
5714static void dropTerminatorIfPresent(Block &block) {
5715 if (!block.empty() && block.back().hasTrait<OpTrait::IsTerminator>()) {
5716 block.back().erase();
5717 }
5718}
5719
5721static void
5722moveBranchWithoutLiftedWrites(Block *srcBlock, Block &destBlock, ArrayRef<IfWriteSlot> slots) {
5723 if (srcBlock) {
5724 for (auto it = srcBlock->begin(), end = srcBlock->end(); it != end;) {
5725 Operation &op = *it++;
5726 if (op.hasTrait<OpTrait::IsTerminator>() || isLiftedWrite(op, slots)) {
5727 continue;
5728 }
5729 op.moveBefore(&destBlock, destBlock.end());
5730 }
5731 }
5732}
5733
5736static void appendYield(
5737 OpBuilder &bldr, Location loc, Block &block, ValueRange priorYieldValues,
5738 ArrayRef<IfWriteSlot> slots, bool isThenBlock, scf::YieldOp originalYield = nullptr
5739) {
5740 SmallVector<Value> yieldValues = llvm::to_vector(priorYieldValues);
5741 llvm::append_range(yieldValues, llvm::map_range(slots, [isThenBlock](const IfWriteSlot &slot) {
5742 WritePodOp writeOp = isThenBlock ? slot.thenWrite : slot.elseWrite;
5743 return writeOp ? writeOp.getValue() : slot.incomingValue;
5744 }));
5745
5746 bldr.setInsertionPointToEnd(&block);
5747 auto newYield = bldr.create<scf::YieldOp>(loc, yieldValues);
5748 if (originalYield) {
5749 preserveDiscardableAttrs(originalYield, newYield);
5750 }
5751}
5752
5758struct LoopPodSlot {
5759 Value podRef;
5760 StringAttr recordName;
5761 Type type;
5762
5764 bool matches(Value findPodRef, StringAttr findRecordName) const {
5765 return this->podRef == findPodRef && this->recordName == findRecordName;
5766 }
5767};
5768
5770static LoopPodSlot *
5771lookupLoopSlot(SmallVectorImpl<LoopPodSlot> &slots, Value podRef, StringAttr recordName) {
5772 auto *it = llvm::find_if(slots, [&podRef, &recordName](const LoopPodSlot &slot) {
5773 return slot.matches(podRef, recordName);
5774 });
5775 return it == slots.end() ? nullptr : &*it;
5776}
5777
5779static bool hasLoopSlot(ArrayRef<LoopPodSlot> slots, Value podRef, StringAttr recordName) {
5780 const auto *it = llvm::find_if(slots, [&podRef, &recordName](const LoopPodSlot &slot) {
5781 return slot.matches(podRef, recordName);
5782 });
5783 return it != slots.end();
5784}
5785
5787static LoopPodSlot &getOrCreateLoopSlot(
5788 SmallVectorImpl<LoopPodSlot> &slots, Value podRef, StringAttr recordName, Type type
5789) {
5790 if (LoopPodSlot *slot = lookupLoopSlot(slots, podRef, recordName)) {
5791 return *slot;
5792 }
5793 slots.push_back(LoopPodSlot {podRef, recordName, type});
5794 return slots.back();
5795}
5796
5798static std::optional<size_t>
5799findLoopSlotIndex(ArrayRef<LoopPodSlot> slots, Value podRef, StringAttr recordName) {
5800 for (auto [idx, slot] : llvm::enumerate(slots)) {
5801 if (slot.podRef == podRef && slot.recordName == recordName) {
5802 return idx;
5803 }
5804 }
5805 return std::nullopt;
5806}
5807
5810static void
5811collectDirectLoopPodSlots(Block &block, Operation *ancestor, SmallVectorImpl<LoopPodSlot> &slots) {
5812 for (Operation &op : block) {
5813 if (auto readOp = dyn_cast<ReadPodOp>(&op)) {
5814 if (!isValueDefinedInside(ancestor, readOp.getPodRef())) {
5815 getOrCreateLoopSlot(
5816 slots, readOp.getPodRef(), readOp.getRecordNameAttr(), readOp.getType()
5817 );
5818 }
5819 continue;
5820 }
5821
5822 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5823 if (!isValueDefinedInside(ancestor, writeOp.getPodRef())) {
5824 getOrCreateLoopSlot(
5825 slots, writeOp.getPodRef(), writeOp.getRecordNameAttr(), writeOp.getValue().getType()
5826 );
5827 }
5828 }
5829 }
5830}
5831
5833static bool opUsesTrackedPodRefDirectly(Operation &op, ArrayRef<LoopPodSlot> slots) {
5834 return llvm::any_of(op.getOperands(), [&slots](Value operand) {
5835 return llvm::any_of(slots, [&operand](const LoopPodSlot &slot) {
5836 return slot.podRef == operand;
5837 });
5838 });
5839}
5840
5842static bool hasNestedTrackedPodAccess(Operation &op, ArrayRef<LoopPodSlot> slots) {
5843 return op
5844 .walk([&op, &slots](Operation *nestedOp) {
5845 if (nestedOp == &op) {
5846 return WalkResult::advance();
5847 }
5848
5849 if (auto readOp = dyn_cast<ReadPodOp>(nestedOp)) {
5850 if (hasLoopSlot(slots, readOp.getPodRef(), readOp.getRecordNameAttr())) {
5851 return WalkResult::interrupt();
5852 }
5853 return WalkResult::advance();
5854 }
5855
5856 if (auto writeOp = dyn_cast<WritePodOp>(nestedOp)) {
5857 if (hasLoopSlot(slots, writeOp.getPodRef(), writeOp.getRecordNameAttr())) {
5858 return WalkResult::interrupt();
5859 }
5860 }
5861 return WalkResult::advance();
5862 }).wasInterrupted();
5863}
5864
5867static bool hasUnliftableLoopPodUses(Block &block, ArrayRef<LoopPodSlot> slots) {
5868 for (Operation &op : block) {
5869 if (isa<ReadPodOp, WritePodOp>(op)) {
5870 continue;
5871 }
5872 if (opUsesTrackedPodRefDirectly(op, slots) || hasNestedTrackedPodAccess(op, slots)) {
5873 return true;
5874 }
5875 }
5876 return false;
5877}
5878
5880template <typename RewriteTerminatorFn>
5881static void cloneLoopBodyWithLiftedPodSlots(
5882 Block &source, PatternRewriter &rewriter, IRMapping &mapping, ArrayRef<LoopPodSlot> slots,
5883 SmallVectorImpl<Value> &slotValues, RewriteTerminatorFn &&rewriteTerminator
5884) {
5885 for (Operation &op : source) {
5886 if (rewriteTerminator(op)) {
5887 continue;
5888 }
5889
5890 if (auto readOp = dyn_cast<ReadPodOp>(&op)) {
5891 if (std::optional<size_t> slotIdx =
5892 findLoopSlotIndex(slots, readOp.getPodRef(), readOp.getRecordNameAttr())) {
5893 mapping.map(readOp.getResult(), slotValues[*slotIdx]);
5894 continue;
5895 }
5896 }
5897
5898 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5899 if (std::optional<size_t> slotIdx =
5900 findLoopSlotIndex(slots, writeOp.getPodRef(), writeOp.getRecordNameAttr())) {
5901 slotValues[*slotIdx] = mapping.lookupOrDefault(writeOp.getValue());
5902 continue;
5903 }
5904 }
5905
5906 rewriter.clone(op, mapping);
5907 }
5908}
5909
5911static void appendIncomingLoopSlotValues(
5912 PatternRewriter &rewriter, Location loc, ArrayRef<LoopPodSlot> slots,
5913 SmallVectorImpl<Value> &values, SmallVectorImpl<Type> *resultTypes = nullptr
5914) {
5915 for (const LoopPodSlot &slot : slots) {
5916 values.push_back(genRead(rewriter, loc, slot.podRef, slot.recordName).getResult());
5917 if (resultTypes) {
5918 resultTypes->push_back(slot.type);
5919 }
5920 }
5921}
5922
5924template <typename ValueRangeLike>
5925static SmallVector<Value>
5926collectTrailingLoopSlotValues(ValueRangeLike values, size_t base, size_t slotCount) {
5927 SmallVector<Value> slotValues;
5928 slotValues.reserve(slotCount);
5929 for (size_t idx = 0; idx < slotCount; ++idx) {
5930 slotValues.push_back(values[llzk::checkedCast<unsigned>(base + idx)]);
5931 }
5932 return slotValues;
5933}
5934
5936static SmallVector<Value>
5937remapValuesAndAppendLoopSlots(ValueRange values, IRMapping &mapping, ValueRange slotValues) {
5938 SmallVector<Value> remappedValues = llvm::map_to_vector(values, [&mapping](Value value) {
5939 return mapping.lookupOrDefault(value);
5940 });
5941 llvm::append_range(remappedValues, slotValues);
5942 return remappedValues;
5943}
5944
5946template <typename GetResultFn>
5947static void writeBackLoopSlotResults(
5948 PatternRewriter &rewriter, Location loc, ArrayRef<LoopPodSlot> slots,
5949 unsigned originalResultCount, GetResultFn &&getResult
5950) {
5951 for (auto [idx, slot] : llvm::enumerate(slots)) {
5952 genWrite(
5953 rewriter, loc, slot.podRef, slot.recordName,
5954 getResult(originalResultCount + llzk::checkedCast<unsigned>(idx))
5955 );
5956 }
5957}
5958
5962class LiftPodWritesFromIfBlocksPattern final : public OpRewritePattern<scf::IfOp> {
5963public:
5964 using OpRewritePattern<scf::IfOp>::OpRewritePattern;
5965
5966 LogicalResult matchAndRewrite(scf::IfOp ifOp, PatternRewriter &rewriter) const override {
5967 SmallVector<IfWriteSlot> slots;
5968 Block &thenBlock = *ifOp.thenBlock();
5969 Block *elseBlock = getElseBlockOrNull(ifOp);
5970 collectDirectWrites(&thenBlock, true, slots);
5971 collectDirectWrites(elseBlock, false, slots);
5972 if (slots.empty()) {
5973 return failure();
5974 }
5975
5976 llvm::erase_if(slots, [&](const IfWriteSlot &slot) {
5977 return isValueDefinedInside(ifOp, slot.podRef) ||
5978 !branchSlotCanBeLifted(&thenBlock, slot.podRef, slot.recordName) ||
5979 !branchSlotCanBeLifted(elseBlock, slot.podRef, slot.recordName);
5980 });
5981 if (slots.empty()) {
5982 return failure();
5983 }
5984
5985 for (IfWriteSlot &slot : slots) {
5986 if (slot.thenWrite && slot.elseWrite) {
5987 continue;
5988 }
5989 rewriter.setInsertionPoint(ifOp);
5990 slot.incomingValue =
5991 genRead(rewriter, ifOp.getLoc(), slot.podRef, slot.recordName).getResult();
5992 }
5993
5994 SmallVector<Type> resultTypes = llvm::to_vector(ifOp.getResultTypes());
5995 llvm::append_range(resultTypes, llvm::map_range(slots, [](auto slot) { return slot.type; }));
5996
5997 scf::YieldOp thenYieldOp = getYieldOp(thenBlock);
5998 SmallVector<Value> originalThenYields;
5999 if (!ifOp.getResults().empty()) {
6000 originalThenYields.append(thenYieldOp.getOperands().begin(), thenYieldOp.getOperands().end());
6001 }
6002
6003 scf::YieldOp elseYieldOp = elseBlock ? getYieldOp(*elseBlock) : nullptr;
6004 SmallVector<Value> originalElseYields;
6005 if (elseBlock && !ifOp.getResults().empty()) {
6006 originalElseYields.append(elseYieldOp.getOperands().begin(), elseYieldOp.getOperands().end());
6007 }
6008
6009 rewriter.setInsertionPoint(ifOp);
6010 auto newIf = rewriter.create<scf::IfOp>(ifOp.getLoc(), resultTypes, ifOp.getCondition(), true);
6011 Block &newThenBlock = *newIf.thenBlock();
6012 Block &newElseBlock = *newIf.elseBlock();
6013 dropTerminatorIfPresent(newThenBlock);
6014 dropTerminatorIfPresent(newElseBlock);
6015
6016 moveBranchWithoutLiftedWrites(&thenBlock, newThenBlock, slots);
6017 moveBranchWithoutLiftedWrites(elseBlock, newElseBlock, slots);
6018 appendYield(
6019 rewriter, ifOp.getLoc(), newThenBlock, originalThenYields, slots, true, thenYieldOp
6020 );
6021 appendYield(
6022 rewriter, ifOp.getLoc(), newElseBlock, originalElseYields, slots, false, elseYieldOp
6023 );
6024
6025 rewriter.setInsertionPointAfter(newIf);
6026 unsigned originalResultCount = ifOp.getNumResults();
6027 for (auto [idx, slot] : llvm::enumerate(slots)) {
6028 genWrite(
6029 rewriter, ifOp.getLoc(), slot.podRef, slot.recordName,
6030 newIf.getResult(originalResultCount + idx)
6031 );
6032 }
6033
6034 rewriter.replaceOp(ifOp, newIf.getResults().take_front(originalResultCount));
6035 return success();
6036 }
6037};
6038
6041class LiftPodAccessesFromForLoopPattern final : public OpRewritePattern<scf::ForOp> {
6042public:
6043 using OpRewritePattern<scf::ForOp>::OpRewritePattern;
6044
6045 LogicalResult matchAndRewrite(scf::ForOp forOp, PatternRewriter &rewriter) const override {
6046 Block &body = *forOp.getBody();
6047 SmallVector<LoopPodSlot> slots;
6048 collectDirectLoopPodSlots(body, forOp.getOperation(), slots);
6049 if (slots.empty() || hasUnliftableLoopPodUses(body, slots)) {
6050 return failure();
6051 }
6052
6053 Location loc = forOp.getLoc();
6054
6055 SmallVector<Value> newInitArgs = llvm::to_vector(forOp.getInitArgs());
6056 rewriter.setInsertionPoint(forOp);
6057 appendIncomingLoopSlotValues(rewriter, loc, slots, newInitArgs);
6058
6059 auto newFor = rewriter.create<scf::ForOp>(
6060 loc, forOp.getLowerBound(), forOp.getUpperBound(), forOp.getStep(), newInitArgs
6061 );
6062 newFor->setAttrs(forOp->getAttrs());
6063
6064 Block &newBody = *newFor.getBody();
6065 dropTerminatorIfPresent(newBody);
6066
6067 IRMapping mapping;
6068 mapping.map(forOp.getInductionVar(), newFor.getInductionVar());
6069 for (auto [idx, oldArg] : llvm::enumerate(forOp.getRegionIterArgs())) {
6070 mapping.map(oldArg, newFor.getRegionIterArg(idx));
6071 }
6072
6073 SmallVector<Value> slotValues = collectTrailingLoopSlotValues(
6074 newFor.getRegionIterArgs(), forOp.getNumRegionIterArgs(), slots.size()
6075 );
6076
6077 rewriter.setInsertionPointToEnd(&newBody);
6078 cloneLoopBodyWithLiftedPodSlots(body, rewriter, mapping, slots, slotValues, [&](Operation &op) {
6079 if (auto yieldOp = dyn_cast<scf::YieldOp>(&op)) {
6080 SmallVector<Value> yieldValues =
6081 remapValuesAndAppendLoopSlots(yieldOp.getOperands(), mapping, slotValues);
6082 preserveDiscardableAttrs(
6083 yieldOp, rewriter.create<scf::YieldOp>(yieldOp.getLoc(), yieldValues)
6084 );
6085 return true;
6086 }
6087 return false;
6088 });
6089
6090 rewriter.setInsertionPointAfter(newFor);
6091 writeBackLoopSlotResults(
6092 rewriter, loc, slots, forOp.getNumResults(),
6093 [&newFor](unsigned resultIdx) { return newFor.getResult(resultIdx); }
6094 );
6095
6096 rewriter.replaceOp(forOp, newFor.getResults().take_front(forOp.getNumResults()));
6097 return success();
6098 }
6099};
6100
6103class LiftPodAccessesFromWhileLoopPattern final : public OpRewritePattern<scf::WhileOp> {
6104public:
6105 using OpRewritePattern<scf::WhileOp>::OpRewritePattern;
6106
6107 LogicalResult matchAndRewrite(scf::WhileOp whileOp, PatternRewriter &rewriter) const override {
6108 Block &beforeBody = *whileOp.getBeforeBody();
6109 Block &afterBody = *whileOp.getAfterBody();
6110
6111 SmallVector<LoopPodSlot> slots;
6112 collectDirectLoopPodSlots(beforeBody, whileOp.getOperation(), slots);
6113 collectDirectLoopPodSlots(afterBody, whileOp.getOperation(), slots);
6114 if (slots.empty() || hasUnliftableLoopPodUses(beforeBody, slots) ||
6115 hasUnliftableLoopPodUses(afterBody, slots)) {
6116 return failure();
6117 }
6118
6119 Location loc = whileOp.getLoc();
6120
6121 SmallVector<Value> newInits = llvm::to_vector(whileOp.getInits());
6122 SmallVector<Type> newResultTypes = llvm::to_vector(whileOp.getResultTypes());
6123 rewriter.setInsertionPoint(whileOp);
6124 appendIncomingLoopSlotValues(rewriter, loc, slots, newInits, &newResultTypes);
6125
6126 auto newWhile = rewriter.create<scf::WhileOp>(loc, newResultTypes, newInits, nullptr, nullptr);
6127 newWhile->setAttrs(whileOp->getAttrs());
6128
6129 Block &newBeforeBody = *newWhile.getBeforeBody();
6130 Block &newAfterBody = *newWhile.getAfterBody();
6131 dropTerminatorIfPresent(newBeforeBody);
6132 dropTerminatorIfPresent(newAfterBody);
6133
6134 IRMapping beforeMapping;
6135 for (auto [oldArg, newArg] : llvm::zip_equal(
6136 whileOp.getBeforeArguments(),
6137 newWhile.getBeforeArguments().take_front(whileOp.getBeforeArguments().size())
6138 )) {
6139 beforeMapping.map(oldArg, newArg);
6140 }
6141
6142 SmallVector<Value> beforeSlotValues = collectTrailingLoopSlotValues(
6143 newWhile.getBeforeArguments(), whileOp.getBeforeArguments().size(), slots.size()
6144 );
6145
6146 rewriter.setInsertionPointToEnd(&newBeforeBody);
6147 cloneLoopBodyWithLiftedPodSlots(
6148 beforeBody, rewriter, beforeMapping, slots, beforeSlotValues, [&](Operation &op) {
6149 if (auto conditionOp = dyn_cast<scf::ConditionOp>(&op)) {
6150 SmallVector<Value> conditionArgs =
6151 remapValuesAndAppendLoopSlots(conditionOp.getArgs(), beforeMapping, beforeSlotValues);
6152 preserveDiscardableAttrs(
6153 conditionOp,
6154 rewriter.create<scf::ConditionOp>(
6155 conditionOp.getLoc(), beforeMapping.lookupOrDefault(conditionOp.getCondition()),
6156 conditionArgs
6157 )
6158 );
6159 return true;
6160 }
6161 return false;
6162 }
6163 );
6164
6165 IRMapping afterMapping;
6166 for (auto [oldArg, newArg] : llvm::zip_equal(
6167 whileOp.getAfterArguments(),
6168 newWhile.getAfterArguments().take_front(whileOp.getAfterArguments().size())
6169 )) {
6170 afterMapping.map(oldArg, newArg);
6171 }
6172
6173 SmallVector<Value> afterSlotValues = collectTrailingLoopSlotValues(
6174 newWhile.getAfterArguments(), whileOp.getAfterArguments().size(), slots.size()
6175 );
6176
6177 rewriter.setInsertionPointToEnd(&newAfterBody);
6178 cloneLoopBodyWithLiftedPodSlots(
6179 afterBody, rewriter, afterMapping, slots, afterSlotValues, [&](Operation &op) {
6180 if (auto yieldOp = dyn_cast<scf::YieldOp>(&op)) {
6181 SmallVector<Value> yieldValues =
6182 remapValuesAndAppendLoopSlots(yieldOp.getOperands(), afterMapping, afterSlotValues);
6183 preserveDiscardableAttrs(
6184 yieldOp, rewriter.create<scf::YieldOp>(yieldOp.getLoc(), yieldValues)
6185 );
6186 return true;
6187 }
6188 return false;
6189 }
6190 );
6191
6192 rewriter.setInsertionPointAfter(newWhile);
6193 writeBackLoopSlotResults(
6194 rewriter, loc, slots, whileOp.getNumResults(),
6195 [&newWhile](unsigned resultIdx) { return newWhile.getResult(resultIdx); }
6196 );
6197
6198 rewriter.replaceOp(whileOp, newWhile.getResults().take_front(whileOp.getNumResults()));
6199 return success();
6200 }
6201};
6202
6208class SplitPodInEmitEqualityPattern final : public OpRewritePattern<constrain::EmitEqualityOp> {
6209 CompatiblePodLeafMaterializationMap &materializedLeaves;
6210
6211public:
6212 SplitPodInEmitEqualityPattern(
6213 MLIRContext *ctx, CompatiblePodLeafMaterializationMap &materializedLeafMap
6214 )
6215 : OpRewritePattern<constrain::EmitEqualityOp>(ctx), materializedLeaves(materializedLeafMap) {}
6216
6217 LogicalResult
6218 matchAndRewrite(constrain::EmitEqualityOp op, PatternRewriter &rewriter) const override {
6219 PodType podTy = getWholePodEqualityType(op);
6220 if (!podTy) {
6221 return failure();
6222 }
6223
6224 return splitWholePodEmitEquality(op, rewriter, materializedLeaves);
6225 }
6226};
6227
6229static LogicalResult
6230applyGreedily(ModuleOp modOp, RewritePatternSet &&patterns, bool *changed = nullptr) {
6231 return applyPatternsGreedily(
6232 modOp->getRegion(0), std::move(patterns),
6233 GreedyRewriteConfig {.fold = false, .cseConstants = false}, changed
6234 );
6235}
6236
6239static LogicalResult step4(ModuleOp modOp) {
6240 CompatiblePodLeafMaterializationMap materializedLeaves;
6241 RewritePatternSet patterns(modOp.getContext());
6242 patterns.add<
6243 FoldReadAfterWriteInBlockPattern, ReplaceIfReadPattern, LiftPodWritesFromIfBlocksPattern,
6244 LiftPodAccessesFromForLoopPattern, LiftPodAccessesFromWhileLoopPattern,
6245 FoldIfCarriedPodReadAfterWritePattern>(patterns.getContext());
6246 patterns.add<SplitPodInEmitEqualityPattern>(patterns.getContext(), materializedLeaves);
6247
6248 LLVM_DEBUG(llvm::dbgs() << "Begin step 4: refactor pod ops within SCF regions\n";);
6249 return applyGreedily(modOp, std::move(patterns));
6250}
6251
6254static bool applyIfCarriedPodReadAfterWritePatterns(ModuleOp modOp) {
6255 RewritePatternSet patterns(modOp.getContext());
6256 patterns.add<FoldIfCarriedPodReadAfterWritePattern>(patterns.getContext());
6257
6258 bool changed = false;
6259 if (failed(applyGreedily(modOp, std::move(patterns), &changed))) {
6260 return false;
6261 }
6262 return changed;
6263}
6264
6267static size_t podTypeScalarizationWeight(Type type) {
6268 auto podTy = dyn_cast<PodType>(type);
6269 if (!podTy) {
6270 return 0;
6271 }
6272
6273 size_t weight = 1;
6274 for (RecordAttr record : podTy.getRecords()) {
6275 weight += podTypeScalarizationWeight(record.getType());
6276 }
6277 return weight;
6278}
6279
6283static size_t podAllocScalarizationWeight(ModuleOp modOp) {
6284 size_t weight = 0;
6285 modOp.walk([&weight](NewPodOp newPodOp) {
6286 weight += podTypeScalarizationWeight(newPodOp.getType());
6287 });
6288 return weight;
6289}
6290
6292inline static bool isResidualPodLikeType(Type type) {
6293 return splittablePod(type) || splittablePodArray(type);
6294}
6295
6297template <typename TypeRangeLike>
6298static size_t countResidualPodLikeTypes(const TypeRangeLike &types) {
6299 size_t count = 0;
6300 for (Type type : types) {
6301 if (isResidualPodLikeType(type)) {
6302 ++count;
6303 }
6304 }
6305 return count;
6306}
6307
6309static bool isResidualPodPlaceholderCast(UnrealizedConversionCastOp castOp) {
6310 return countResidualPodLikeTypes(castOp.getOperandTypes()) != 0 ||
6311 countResidualPodLikeTypes(castOp.getResultTypes()) != 0;
6312}
6313
6315static size_t countResidualPodIR(ModuleOp modOp) {
6316 size_t count = 0;
6317 modOp.walk([&count](Operation *op) {
6318 if (isa<NewPodOp, ReadPodOp, WritePodOp>(op)) {
6319 ++count;
6320 } else if (auto castOp = dyn_cast<UnrealizedConversionCastOp>(op);
6321 castOp && isResidualPodPlaceholderCast(castOp)) {
6322 ++count;
6323 }
6324
6325 count += countResidualPodLikeTypes(op->getOperandTypes());
6326 count += countResidualPodLikeTypes(op->getResultTypes());
6327
6328 for (Region &region : op->getRegions()) {
6329 for (Block &block : region) {
6330 count += countResidualPodLikeTypes(block.getArgumentTypes());
6331 }
6332 }
6333
6334 if (auto funcDef = dyn_cast<FuncDefOp>(op)) {
6335 FunctionType funcTy = funcDef.getFunctionType();
6336 count += countResidualPodLikeTypes(funcTy.getInputs());
6337 count += countResidualPodLikeTypes(funcTy.getResults());
6338 } else if (auto memberDef = dyn_cast<MemberDefOp>(op)) {
6339 count += isResidualPodLikeType(memberDef.getType()) ? 1 : 0;
6340 }
6341 });
6342 return count;
6343}
6344
6346static LogicalResult rejectRaggedNestedLeafArrayLengths(ModuleOp modOp) {
6347 WalkResult result = modOp.walk([](ArrayLengthOp op) {
6348 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getArrRef());
6349 if (raggedKind.empty()) {
6350 return WalkResult::advance();
6351 }
6352 op.emitOpError() << "cannot lower nested " << raggedKind
6353 << " array leaf length after reading an array-of-POD element without "
6354 "per-element shape witnesses";
6355 return WalkResult::interrupt();
6356 });
6357 return failure(result.wasInterrupted());
6358}
6359
6361static LogicalResult rejectRaggedNestedLeafArrayEqualities(ModuleOp modOp) {
6362 WalkResult result = modOp.walk([](constrain::EmitEqualityOp op) {
6363 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getLhs());
6364 if (raggedKind.empty()) {
6365 raggedKind = getTaggedRaggedNestedLeafKind(op.getRhs());
6366 }
6367 if (raggedKind.empty()) {
6368 return WalkResult::advance();
6369 }
6370 op.emitOpError() << "cannot lower nested " << raggedKind
6371 << " array leaf equality after reading an array-of-POD element without "
6372 "per-element shape witnesses";
6373 return WalkResult::interrupt();
6374 });
6375 return failure(result.wasInterrupted());
6376}
6377
6379static LogicalResult rejectRaggedNestedLeafArrayContainments(ModuleOp modOp) {
6380 WalkResult result = modOp.walk([](constrain::EmitContainmentOp op) {
6381 if (succeeded(rejectRaggedNestedLeafContainment(op))) {
6382 return WalkResult::advance();
6383 }
6384 return WalkResult::interrupt();
6385 });
6386 return failure(result.wasInterrupted());
6387}
6388
6390static LogicalResult rejectRaggedNestedLeafBoundaryCrossings(ModuleOp modOp) {
6391 WalkResult result = modOp.walk([](Operation *op) {
6392 LogicalResult check = TypeSwitch<Operation *, LogicalResult>(op)
6393 .Case<CallOp>([](CallOp callOp) {
6394 return rejectRaggedNestedLeafBoundaryCrossing(
6395 callOp.getOperation(), callOp.getArgOperands(), "across a function call"
6396 );
6397 })
6398 .Case<ReturnOp>([](ReturnOp returnOp) {
6399 return rejectRaggedNestedLeafBoundaryCrossing(
6400 returnOp.getOperation(), returnOp.getOperands(), "across a function return"
6401 );
6402 })
6403 .Case<scf::ForOp>([](scf::ForOp forOp) {
6404 return rejectRaggedNestedLeafBoundaryCrossing(
6405 forOp.getOperation(), forOp.getInitArgs(), "into an scf.for region"
6406 );
6407 })
6408 .Case<scf::WhileOp>([](scf::WhileOp whileOp) {
6409 return rejectRaggedNestedLeafBoundaryCrossing(
6410 whileOp.getOperation(), whileOp.getInits(), "into an scf.while region"
6411 );
6412 })
6413 .Case<scf::ConditionOp>([](scf::ConditionOp conditionOp) {
6414 return rejectRaggedNestedLeafBoundaryCrossing(
6415 conditionOp.getOperation(), conditionOp.getArgs(), "across an scf.condition boundary"
6416 );
6417 })
6418 .Case<scf::YieldOp>([](scf::YieldOp yieldOp) {
6419 return rejectRaggedNestedLeafBoundaryCrossing(
6420 yieldOp.getOperation(), yieldOp.getOperands(), "across an scf.yield boundary"
6421 );
6422 }).Default([](Operation *) { return success(); });
6423 return failed(check) ? WalkResult::interrupt() : WalkResult::advance();
6424 });
6425 return failure(result.wasInterrupted());
6426}
6427
6429static LogicalResult rejectRemainingRaggedNestedLeafUses(ModuleOp modOp) {
6430 if (failed(rejectRaggedNestedLeafArrayLengths(modOp))) {
6431 return failure();
6432 }
6433 if (failed(rejectRaggedNestedLeafArrayEqualities(modOp))) {
6434 return failure();
6435 }
6436 if (failed(rejectRaggedNestedLeafArrayContainments(modOp))) {
6437 return failure();
6438 }
6439 return rejectRaggedNestedLeafBoundaryCrossings(modOp);
6440}
6441
6443class PassImpl : public llzk::pod::impl::PodToScalarPassBase<PassImpl> {
6444 using Base = PodToScalarPassBase<PassImpl>;
6445 using Base::Base;
6446
6447 LogicalResult runScalarizeAndCleanupPipeline(ModuleOp module) {
6448 // 1. Use SROA (Destructurable* interfaces) to split each pod with `N` records into `N` pods
6449 // with 1 record each. This is necessary because the mem2reg pass cannot deal with splitting
6450 // up memory, i.e., it can only convert scalar memory access into SSA values.
6451 // 2. The mem2reg pass converts the size 1 pod allocations and accesses into SSA values.
6452 OpPassManager scalarizePM(ModuleOp::getOperationName());
6453 scalarizePM.addPass(createSpecializedSROAPass<NewPodOp>());
6454 scalarizePM.addPass(createSpecializedMem2RegPass<NewPodOp>());
6455
6456 // Cleanup allocations made dead by memory promotion and other dead SSA values.
6457 OpPassManager cleanupPM(ModuleOp::getOperationName());
6459 RemoveUnusedDiscardableAllocationsPassOptions {
6460 .allocatorOpName = CreateArrayOp::getOperationName().str()
6461 }
6462 ));
6464 RemoveUnusedDiscardableAllocationsPassOptions {
6465 .allocatorOpName = NewPodOp::getOperationName().str()
6466 }
6467 ));
6468 cleanupPM.addPass(createRemoveDeadValuesWorkaroundPass());
6469
6470 size_t podAllocWeight = podAllocScalarizationWeight(module);
6471 while (podAllocWeight != 0) {
6472 if (failed(runPipeline(scalarizePM, module))) {
6473 return failure();
6474 }
6475
6476 // SROA+mem2reg can expose `scf.if`-carried POD values that become redundant after a
6477 // same-record write from another `scf.if` result. Fold those reads and clean up before
6478 // checking convergence.
6479 bool foldedIfCarriedRead = applyIfCarriedPodReadAfterWritePatterns(module);
6480 if (failed(runPipeline(cleanupPM, module))) {
6481 return failure();
6482 }
6483
6484 // Nested PODs can become visible only after an outer single-record POD has been promoted,
6485 // and SROA can transiently increase allocation count while splitting aggregates. Keep
6486 // iterating until the allocation-weight heuristic reaches a fixed point.
6487 size_t nextPodAllocWeight = podAllocScalarizationWeight(module);
6488 if (!foldedIfCarriedRead && nextPodAllocWeight == podAllocWeight) {
6489 break;
6490 }
6491 podAllocWeight = nextPodAllocWeight;
6492 }
6493
6494 return success();
6495 }
6496
6497 void runOnOperation() override {
6498 ModuleOp module = getOperation();
6499 if (failed(step0(module))) {
6500 return signalPassFailure();
6501 }
6502 LLVM_DEBUG({
6503 llvm::dbgs() << "After step 0:\n";
6504 module.dump();
6505 });
6506
6507 size_t previousResidualCount = std::numeric_limits<size_t>::max();
6508 while (true) {
6509 // This is divided into 2 steps to simplify the implementation for member-related ops. The
6510 // issue is that the conversions for member read/write expect the mapping of record name to
6511 // member name+type to already be populated for the referenced member (although this could be
6512 // computed on demand if desired but it complicates the implementation a bit).
6513 SymbolTableCollection symTables;
6514 MemberReplacementMap memberRepMap;
6515 if (failed(step1(module, symTables, memberRepMap))) {
6516 return signalPassFailure();
6517 }
6518 LLVM_DEBUG({
6519 llvm::dbgs() << "After step 1:\n";
6520 module.dump();
6521 });
6522
6523 if (failed(step2(module, symTables, memberRepMap))) {
6524 return signalPassFailure();
6525 }
6526 LLVM_DEBUG({
6527 llvm::dbgs() << "After step 2:\n";
6528 module.dump();
6529 });
6530
6531 if (failed(step3(module, symTables, memberRepMap))) {
6532 return signalPassFailure();
6533 }
6534 LLVM_DEBUG({
6535 llvm::dbgs() << "After step 3:\n";
6536 module.dump();
6537 });
6538
6539 if (failed(rejectRemainingRaggedNestedLeafUses(module))) {
6540 signalPassFailure();
6541 return;
6542 }
6543
6544 if (failed(step4(module))) {
6545 return signalPassFailure();
6546 }
6547 LLVM_DEBUG({
6548 llvm::dbgs() << "After step 4:\n";
6549 module.dump();
6550 });
6551
6552 if (failed(runScalarizeAndCleanupPipeline(module))) {
6553 signalPassFailure();
6554 return;
6555 }
6556 LLVM_DEBUG({
6557 llvm::dbgs() << "After SROA+Mem2Reg pipeline:\n";
6558 module.dump();
6559 });
6560
6561 if (failed(rejectRemainingRaggedNestedLeafUses(module))) {
6562 signalPassFailure();
6563 return;
6564 }
6565
6566 size_t residualCount = countResidualPodIR(module);
6567 if (residualCount == 0) {
6568 break;
6569 }
6570 if (residualCount >= previousResidualCount) {
6571 std::string residualIR;
6572 llvm::raw_string_ostream os(residualIR);
6573 module.print(os);
6574 module.emitError() << "llzk-pod-to-scalar left residual pod IR after reaching a fixpoint ("
6575 << residualCount << " residual pod-like items remain)\n"
6576 << residualIR;
6577 signalPassFailure();
6578 return;
6579 }
6580 previousResidualCount = residualCount;
6581 }
6582 }
6583};
6584
6585} // namespace
within a display generated by the Derivative if and wherever such third party notices normally appear The contents of the NOTICE file are for informational purposes only and do not modify the License You may add Your own attribution notices within Derivative Works that You alongside or as an addendum to the NOTICE text from the provided that such additional attribution notices cannot be construed as modifying the License You may add Your own copyright statement to Your modifications and may provide additional or different license terms and conditions for or distribution of Your or for any such Derivative Works as a provided Your and distribution of the Work otherwise complies with the conditions stated in this License Submission of Contributions Unless You explicitly state any Contribution intentionally submitted for inclusion in the Work by You to the Licensor shall be under the terms and conditions of this without any additional terms or conditions Notwithstanding the nothing herein shall supersede or modify the terms of any separate license agreement you may have executed with Licensor regarding such Contributions Trademarks This License does not grant permission to use the trade names
Definition LICENSE.txt:139
Apache License January AND DISTRIBUTION Definitions License shall mean the terms and conditions for and distribution as defined by Sections through of this document Licensor shall mean the copyright owner or entity authorized by the copyright owner that is granting the License Legal Entity shall mean the union of the acting entity and all other entities that control are controlled by or are under common control with that entity For the purposes of this definition control direct or to cause the direction or management of such whether by contract or including but not limited to software source documentation source
Definition LICENSE.txt:28
Apache License January AND DISTRIBUTION Definitions License shall mean the terms and conditions for use
Definition LICENSE.txt:9
#define check(x)
Definition Ops.cpp:286
Provides SpecializedSROA<AllocOpTy> and SpecializedMem2Reg<AllocOpTy>: pass templates that replicate ...
::mlir::Value genRead(::mlir::OpBuilder &bldr, ::mlir::Location loc, ::mlir::Value arrayRef, ::mlir::ValueRange indices)
Create an array.read or array.extract for one concrete element or subarray.
static ::llvm::SmallVector<::mlir::Value > genIndexConstants(::mlir::OpBuilder &bldr, ::mlir::Location loc, ::mlir::ArrayAttr index)
Generate arith.constant indices for one static array element position.
Definition Ops.cpp:228
inline ::llzk::array::ArrayType getArrRefType()
Gets the type of the referenced base array.
Definition Ops.h.inc:192
::mlir::TypedValue<::llzk::array::ArrayType > getArrRef()
Definition Ops.h.inc:146
::mlir::TypedValue<::mlir::IndexType > getDim()
Definition Ops.h.inc:150
ArrayType cloneWith(std::optional<::llvm::ArrayRef< int64_t > > shape, ::mlir::Type elementType) const
Clone this type with the given shape and element type.
std::optional<::llvm::SmallVector<::mlir::ArrayAttr > > getSubelementIndices() const
Return a list of all valid indices for this ArrayType.
Definition Types.cpp:113
::mlir::Type getElementType() const
::mlir::Type getSelectionType(size_t numIndices) const
Return the type produced by selecting/removing numIndices leading dimensions.
Definition Types.cpp:145
static ArrayType get(::mlir::Type elementType, ::llvm::ArrayRef<::mlir::Attribute > dimensionSizes)
Definition Types.cpp.inc:83
::llvm::ArrayRef<::mlir::Attribute > getDimensionSizes() const
::mlir::TypedValue<::llzk::array::ArrayType > getResult()
Definition Ops.h.inc:408
::mlir::DenseI32ArrayAttr getNumDimsPerMapAttr()
Definition Ops.h.inc:421
static constexpr ::llvm::StringLiteral getOperationName()
Definition Ops.h.inc:377
::mlir::Operation::operand_range getElements()
Definition Ops.h.inc:388
::mlir::TypedValue<::llzk::array::ArrayType > getResult()
Definition Ops.h.inc:613
::mlir::TypedValue<::llzk::array::ArrayType > getArrRef()
Definition Ops.h.inc:589
::mlir::TypedValue<::llzk::array::ArrayType > getRvalue()
Definition Ops.h.inc:757
::mlir::TypedValue<::llzk::array::ArrayType > getArrRef()
Definition Ops.h.inc:749
inline ::llzk::array::ArrayType getArrRefType()
Gets the type of the referenced base array.
Definition Ops.h.inc:961
::mlir::TypedValue<::llzk::array::ArrayType > getArrRef()
Definition Ops.h.inc:899
inline ::llzk::array::ArrayType getArrRefType()
Gets the type of the referenced base array.
Definition Ops.h.inc:1125
bool hasPublicAttr()
Returns whether this member is a public output.
Definition Ops.h.inc:463
void setPublicAttr(bool newValue=true)
Adds or removes the unit llzk.pub attribute according to newValue.
Definition Ops.cpp:566
::mlir::Attribute removeSignalAttr()
Definition Ops.h.inc:432
::mlir::StringAttr getSymNameAttr()
Definition Ops.h.inc:386
void setType(::mlir::Type attrValue)
Definition Ops.cpp.inc:556
::mlir::TypedValue<::llzk::component::StructType > getComponent()
Definition Ops.h.inc:691
::std::optional<::mlir::Attribute > getTableOffset()
Definition Ops.cpp.inc:979
::mlir::FlatSymbolRefAttr getMemberNameAttr()
Definition Ops.h.inc:728
::llvm::ArrayRef< int32_t > getNumDimsPerMap()
Definition Ops.cpp.inc:984
::mlir::TypedValue<::llzk::component::StructType > getComponent()
Definition Ops.h.inc:956
::mlir::FlatSymbolRefAttr getMemberNameAttr()
Definition Ops.h.inc:993
::mlir::TypedValue<::mlir::Type > getVal()
Definition Ops.h.inc:960
::mlir::FailureOr< SymbolLookupResult< StructDefOp > > getDefinition(::mlir::SymbolTableCollection &symbolTable, ::mlir::Operation *op, bool reportMissing=true) const
Gets the struct op that defines this struct.
Definition Types.cpp:26
::mlir::ArrayAttr getParams() const
::mlir::TypedValue<::mlir::Type > getRhs()
Definition Ops.h.inc:130
::mlir::TypedValue<::llzk::array::ArrayType > getLhs()
Definition Ops.h.inc:126
::mlir::TypedValue<::mlir::Type > getLhs()
Definition Ops.h.inc:272
::mlir::TypedValue<::mlir::Type > getRhs()
Definition Ops.h.inc:276
::llvm::SmallVector< RangeT > getMapOperands()
Definition Ops.h.inc:175
::mlir::MutableOperandRange getArgOperandsMutable()
Definition Ops.cpp.inc:223
::mlir::Operation::operand_range getArgOperands()
Definition Ops.h.inc:266
CallOpAdaptor Adaptor
Definition Ops.h.inc:205
::mlir::FailureOr<::llzk::SymbolLookupResult<::llzk::function::FuncDefOp > > getCalleeTarget(::mlir::SymbolTableCollection &tables)
Resolve and return the target FuncDefOp for this CallOp.
Definition Ops.cpp:1203
::mlir::FunctionType getFunctionType()
Definition Ops.cpp.inc:984
void setArgAttrsAttr(::mlir::ArrayAttr attr)
Definition Ops.h.inc:741
::llvm::ArrayRef<::mlir::Type > getArgumentTypes()
Required by FunctionOpInterface.
Definition Ops.h.inc:870
void setResAttrsAttr(::mlir::ArrayAttr attr)
Definition Ops.h.inc:745
::mlir::ArrayAttr getArgAttrsAttr()
Definition Ops.h.inc:721
void setFunctionType(::mlir::FunctionType attrValue)
Definition Ops.cpp.inc:1003
::llvm::ArrayRef<::mlir::Type > getResultTypes()
Required by FunctionOpInterface.
Definition Ops.h.inc:874
::mlir::Region & getBody()
Definition Ops.h.inc:698
::mlir::ArrayAttr getResAttrsAttr()
Definition Ops.h.inc:726
::mlir::MutableOperandRange getOperandsMutable()
Definition Ops.cpp.inc:1169
::mlir::Operation::operand_range getOperands()
Definition Ops.h.inc:1011
::mlir::Operation::operand_range getInitialValues()
Definition Ops.h.inc:237
::mlir::OperandRangeRange getMapOperands()
Definition Ops.h.inc:241
::llvm::ArrayRef< int32_t > getNumDimsPerMap()
Definition Ops.cpp.inc:440
::mlir::TypedValue<::llzk::pod::PodType > getResult()
Definition Ops.h.inc:257
void setInitializedRecordsAttr(::mlir::ArrayAttr attr)
Definition Ops.h.inc:285
::mlir::MutableOperandRange getInitialValuesMutable()
Definition Ops.cpp.inc:185
::mlir::ArrayAttr getInitializedRecords()
Definition Ops.cpp.inc:435
static constexpr ::llvm::StringLiteral getOperationName()
Definition Ops.h.inc:226
::llvm::ArrayRef<::llzk::pod::RecordAttr > getRecords() const
::llvm::StringRef getRecordName()
Definition Ops.cpp.inc:632
::mlir::TypedValue<::mlir::Type > getResult()
Definition Ops.h.inc:499
inline ::llzk::pod::PodType getPodRefType()
Gets the type of the referenced pod.
Definition Ops.h.inc:558
::mlir::TypedValue<::llzk::pod::PodType > getPodRef()
Definition Ops.h.inc:480
::mlir::StringAttr getRecordNameAttr()
Definition Ops.h.inc:512
inline ::llzk::pod::PodType getPodRefType()
Gets the type of the referenced pod.
Definition Ops.h.inc:795
::mlir::StringAttr getRecordNameAttr()
Definition Ops.h.inc:749
::mlir::TypedValue<::mlir::Type > getValue()
Definition Ops.h.inc:716
::mlir::TypedValue<::llzk::pod::PodType > getPodRef()
Definition Ops.h.inc:712
::llvm::StringRef getRecordName()
Definition Ops.cpp.inc:977
::mlir::TypedValue<::mlir::Type > getInput()
Definition Ops.h.inc:1327
std::unique_ptr<::mlir::Pass > createLowerBoolQuantifiersPass()
constexpr char ARG_NAME_ATTR_NAME[]
Attribute name for source-level function argument names.
Definition Ops.h:35
constexpr char RES_NAME_ATTR_NAME[]
Attribute name for source-level function result names.
Definition Ops.h:38
bool appendValuesWithExactTypes(mlir::ValueRange values, mlir::TypeRange expectedTypes, llvm::SmallVectorImpl< mlir::Value > &output)
Append values only when every entry exactly matches the corresponding expected type.
mlir::Operation * findNearestLoopCarriedPodAccess(ReadPodOp readOp)
Return the nearest enclosing SCF loop that carries writes to the same external POD record.
std::unique_ptr<::mlir::Pass > createTypeVarInferencePass()
std::unique_ptr<::mlir::Pass > createRemoveUnusedDiscardableAllocationsPass()
std::unique_ptr< mlir::Pass > createRemoveDeadValuesWorkaroundPass()
mlir::ArrayAttr replicateFunctionNameAttrsAsNeeded(mlir::ArrayAttr origAttrs, const llvm::SmallVector< size_t > &originalIdxToSize, const llvm::SmallVector< mlir::Type > &newTypes, llvm::StringRef functionNameAttrName, llvm::ArrayRef< std::optional< llvm::StringRef > > origNames={}, llvm::ArrayRef< llvm::StringRef > existingNames={}, llvm::ArrayRef< llvm::SmallVector< std::string > > splitNameSuffixes={})
Expand function arg/result attribute arrays to match a split signature, rewriting name attrs with the...
bool isNullOrEmpty(mlir::ArrayAttr a)
mlir::DictionaryAttr withFunctionArgNameAttr(mlir::DictionaryAttr attrs, llvm::StringRef name)
Return a copy of the given argument attribute dictionary with function.arg_name set to name.
constexpr T checkedCast(U u) noexcept
Definition Compare.h:81
ArrayType flattenArrayElementType(ArrayType outerArrTy, Type elementType)
std::unique_ptr< SpecializedMem2Reg< AllocOpTy > > createSpecializedMem2RegPass()
bool isDynamic(IntegerAttr intAttr)
bool typesUnify(Type lhs, Type rhs, ArrayRef< StringRef > rhsReversePrefix, UnificationMap *unifications)
function::CallOp createCallPreservingInstantiationOperands(mlir::Location loc, mlir::TypeRange newResultTypes, function::CallOp oldCall, llvm::ArrayRef< mlir::ValueRange > mapOperands, mlir::ValueRange argOperands, mlir::ConversionPatternRewriter &rewriter)
Rebuild a function.call while preserving explicit instantiation state from oldCall.
SplitFunctionNameInfo collectSplitFunctionNameInfo(mlir::ArrayRef< mlir::Type > origTypes, GetNameAttrFn &&getNameAttr, GetSplitSuffixesFn &&getSplitSuffixes)
Collect function arg/result names and split suffixes from a list of original types.
bool hasParentThatIsa(mlir::Operation *op)
Return true if the parameter has a parent/ancestor op that is an instance of one of the template type...
Definition OpHelpers.h:62
bool hasAffineMapAttr(Type type)
mlir::Operation * create(MlirOpBuilder cBuilder, MlirLocation cLocation, Args &&...args)
Creates a new operation using an ODS build method.
Definition Builder.h:41
std::unique_ptr< SpecializedSROA< AllocOpTy > > createSpecializedSROAPass()
std::string reserveUniqueAttrName(llvm::StringSet<> &usedNames, llvm::StringRef desiredName)
Reserve and return a unique function argument/result name based on desiredName.
static unsigned getHashValue(const CompatiblePodLeafMaterializationKey &key)
static CompatiblePodLeafMaterializationKey getTombstoneKey()
static bool isEqual(const CompatiblePodLeafMaterializationKey &lhs, const CompatiblePodLeafMaterializationKey &rhs)
static CompatiblePodLeafMaterializationKey getEmptyKey()
static unsigned getHashValue(const RecordChain &chain)
static bool isEqual(const RecordChain &lhs, const RecordChain &rhs)
llvm::SmallVector< std::optional< llvm::StringRef > > originalNames
llvm::SmallVector< llvm::StringRef > existingNames
llvm::SmallVector< llvm::SmallVector< std::string > > splitNameSuffixes