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 walkContains<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 walkContains<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 walkContains<ReadPodOp>(op, [&podRef, &recordName](ReadPodOp readOp) {
288 return isSamePodRecord(readOp, podRef, recordName);
289 });
290}
291
293static bool hasValueUse(Operation &op, Value value) {
294 return walkContains<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(const 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 matchAndRewrite(
2181 MemberDefOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
2182 ) const override {
2183 if (legal(op)) {
2184 return failure();
2185 }
2186 StructDefOp inStruct = op->getParentOfType<StructDefOp>();
2187 assert(inStruct);
2188 LocalMemberReplacementMap &localRepMapRef = repMapRef[inStruct][op.getSymNameAttr()];
2189
2190 PodType podTy = llvm::cast<PodType>(adaptor.getType()); // safe per legal() check
2191
2192 SymbolTable &structSymbolTable = tables.getSymbolTable(inStruct);
2193 SmallVector<StringAttr> recordChain;
2194 flattenPodMemberIntoLeaves(op, podTy, recordChain, localRepMapRef, structSymbolTable, rewriter);
2195 rewriter.eraseOp(op);
2196 return success();
2197 }
2198};
2199
2201class SplitPodArrayInMemberDefOp : public OpConversionPattern<MemberDefOp> {
2202 SymbolTableCollection &tables;
2203 MemberReplacementMap &repMapRef;
2204
2205public:
2206 SplitPodArrayInMemberDefOp(
2207 MLIRContext *ctx, SymbolTableCollection &symTables, MemberReplacementMap &memberRepMap
2208 )
2209 : OpConversionPattern<MemberDefOp>(ctx), tables(symTables), repMapRef(memberRepMap) {}
2210
2211 inline static bool legal(MemberDefOp op) { return !splittablePodArray(op.getType()); }
2212
2213 LogicalResult matchAndRewrite(
2214 MemberDefOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
2215 ) const override {
2216 if (legal(op)) {
2217 return failure();
2218 }
2219 StructDefOp inStruct = op->getParentOfType<StructDefOp>();
2220 assert(inStruct);
2221 LocalMemberReplacementMap &localRepMapRef = repMapRef[inStruct][op.getSymNameAttr()];
2222
2223 ArrayType arrTy = llvm::cast<ArrayType>(adaptor.getType());
2224 SmallVector<RecordChain> splitIds;
2225 SmallVector<Type> splitTypes;
2226 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2227 if (splitTypes.empty()) {
2228 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2229 rewriter.modifyOpInPlace(op, [&]() {
2230 op.setType(carrierTy);
2231 op.removeSignalAttr();
2232 });
2233 localRepMapRef[RecordChain()] = std::make_pair(op.getSymNameAttr(), carrierTy);
2234 return success();
2235 }
2236
2237 SymbolTable &structSymbolTable = tables.getSymbolTable(inStruct);
2238 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2239 StringAttr name = id.getFlattenedMemberName(op.getContext(), op.getSymNameAttr());
2240 MemberDefOp newMember = rewriter.create<MemberDefOp>(
2241 op.getLoc(), name, splitType, op.getSignal(), op.getColumn()
2242 );
2243 preserveDiscardableAttrs(op, newMember);
2244 newMember.setPublicAttr(op.hasPublicAttr());
2245 localRepMapRef[id] = std::make_pair(structSymbolTable.insert(newMember), splitType);
2246 }
2247 if (needsPodArrayShapeCarrier(arrTy)) {
2248 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2249 StringAttr carrierName =
2250 getSplitPodArrayShapeMemberName(op.getContext(), op.getSymNameAttr());
2251 MemberDefOp carrierMember =
2252 rewriter.create<MemberDefOp>(op.getLoc(), carrierName, carrierTy, false, op.getColumn());
2253 preserveDiscardableAttrs(op, carrierMember);
2254 carrierMember.setPublicAttr(op.hasPublicAttr());
2255 localRepMapRef[RecordChain()] =
2256 std::make_pair(structSymbolTable.insert(carrierMember), carrierTy);
2257 }
2258 rewriter.eraseOp(op);
2259 return success();
2260 }
2261};
2262
2265static LogicalResult
2266step1(ModuleOp modOp, SymbolTableCollection &symTables, MemberReplacementMap &memberRepMap) {
2267 MLIRContext *ctx = modOp.getContext();
2268
2269 RewritePatternSet patterns(ctx);
2270
2271 patterns.add<SplitPodInMemberDefOp, SplitPodArrayInMemberDefOp>(ctx, symTables, memberRepMap);
2272
2273 ConversionTarget target(*ctx);
2274 baseTargetSetup(target);
2275 target.addLegalOp<UnrealizedConversionCastOp>();
2276 target.addDynamicallyLegalOp<MemberDefOp>([](MemberDefOp op) {
2277 return SplitPodInMemberDefOp::legal(op) && SplitPodArrayInMemberDefOp::legal(op);
2278 });
2279
2280 LLVM_DEBUG(llvm::dbgs() << "Begin step 1: split pod-type and array-of-pod members\n";);
2281 return applyFullConversion(modOp, target, std::move(patterns));
2282}
2283
2289class PodArrayTypeConverter : public TypeConverter {
2290public:
2291 PodArrayTypeConverter() {
2292 addConversion([](Type type) { return type; });
2293 addConversion(
2294 [](ArrayType arrTy, SmallVectorImpl<Type> &results) -> std::optional<LogicalResult> {
2295 if (!splittablePodArray(arrTy)) {
2296 return std::nullopt;
2297 }
2298 convertPodArrayTypeTo(arrTy, results);
2299 return success();
2300 }
2301 );
2302
2303 auto materializeCast = [](OpBuilder &bldr, Type targetType, ValueRange inputs,
2304 Location loc) -> Value {
2305 if (inputs.size() != 1 || !typesUnify(inputs.front().getType(), targetType)) {
2306 return {};
2307 }
2308 return castValueToTypeIfNeeded(bldr, loc, inputs.front(), targetType);
2309 };
2310 addTargetMaterialization(materializeCast);
2311 addArgumentMaterialization(materializeCast);
2312 addSourceMaterialization(materializeCast);
2313 }
2314};
2315
2317class SplitPodArrayNonDetOp : public OpConversionPattern<NonDetOp> {
2318public:
2319 using OpConversionPattern<NonDetOp>::OpConversionPattern;
2320
2321 static bool legal(NonDetOp op) { return !splittablePodArray(op.getType()); }
2322
2323 LogicalResult
2324 matchAndRewrite(NonDetOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
2325 if (legal(op)) {
2326 return failure();
2327 }
2328 SmallVector<Type> splitTypes;
2329 splitPodArrayTypeTo(op.getType(), splitTypes);
2330 if (splitTypes.empty()) {
2331 preserveDiscardableAttrs(
2332 op, rewriter.replaceOpWithNewOp<NonDetOp>(
2333 op, getPodArrayShapeCarrierType(llvm::cast<ArrayType>(op.getType()))
2334 )
2335 );
2336 return success();
2337 }
2338 SmallVector<Value> replacements;
2339 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
2340 replacements.reserve(splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
2341 for (Type splitType : splitTypes) {
2342 replacements.push_back(
2343 preserveDiscardableAttrs(op, rewriter.create<NonDetOp>(op.getLoc(), splitType))
2344 );
2345 }
2346 if (needsPodArrayShapeCarrier(arrTy)) {
2347 replacements.push_back(preserveDiscardableAttrs(
2348 op, rewriter.create<NonDetOp>(op.getLoc(), getPodArrayShapeCarrierType(arrTy))
2349 ));
2350 }
2351 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2352 return success();
2353 }
2354};
2355
2366class SplitPodArrayCreateArrayOp : public OpConversionPattern<CreateArrayOp> {
2367public:
2368 using OpConversionPattern<CreateArrayOp>::OpConversionPattern;
2369
2370 static bool legal(CreateArrayOp op) { return !splittablePodArray(op.getType()); }
2371
2372 LogicalResult matchAndRewrite(
2373 CreateArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2374 ) const override {
2375 if (legal(op)) {
2376 return failure();
2377 }
2378 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
2379 SmallVector<RecordChain> splitIds;
2380 SmallVector<Type> splitTypes;
2381 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2382 if (splitTypes.empty()) {
2383 ArrayType carrierTy = getPodArrayShapeCarrierType(arrTy);
2384 if (adaptor.getMapOperands().empty()) {
2385 preserveDiscardableAttrs(op, rewriter.replaceOpWithNewOp<CreateArrayOp>(op, carrierTy));
2386 return success();
2387 }
2388
2389 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2390 preserveDiscardableAttrs(
2391 op, rewriter.replaceOpWithNewOp<CreateArrayOp>(
2392 op, carrierTy, mapOperands.ranges, op.getNumDimsPerMapAttr()
2393 )
2394 );
2395 return success();
2396 }
2397
2398 SmallVector<Value> replacements;
2399 replacements.reserve(splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
2400 DenseI32ArrayAttr numDimsPerMap = op.getNumDimsPerMapAttr();
2401 if (isNullOrEmpty(numDimsPerMap)) {
2402 if (adaptor.getElements().empty()) {
2403 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2404 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2405 replacements.push_back(
2406 createSplitPodArrayReplacement(op, op.getLoc(), arrTy, id, preciseSplitType, rewriter)
2407 );
2408 }
2409 if (needsPodArrayShapeCarrier(arrTy)) {
2410 Value shapeCarrier =
2411 materializeArrayLengthCarrier(op.getResult(), arrTy, op.getLoc(), rewriter);
2412 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
2413 replacements.push_back(shapeCarrier);
2414 }
2415 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2416 return success();
2417 }
2418
2419 auto elementIndices = arrTy.getSubelementIndices();
2420 assert(elementIndices && "array.new with explicit elements requires a static array shape");
2421 assert(
2422 elementIndices->size() == adaptor.getElements().size() &&
2423 "array.new element count must match the outer array cardinality"
2424 );
2425
2426 // Inline initializers are linearized only across the original outer array dimensions. When
2427 // a flattened POD leaf is itself an array, populate the rewritten split array one outer
2428 // element at a time so each leaf array becomes a subarray insert rather than a malformed
2429 // inline operand to the flattened `array.new`.
2430 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2431 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2432 ArrayType storageSplitType = getSplitPodArrayStorageType(arrTy, id.nameList);
2433
2434 SmallVector<Value> leafValues;
2435 leafValues.reserve(adaptor.getElements().size());
2436 for (ValueRange elementRange : adaptor.getElements()) {
2437 Value element = getSingleConvertedValue(elementRange);
2438 leafValues.push_back(genReadAlongPath(rewriter, op.getLoc(), element, id));
2439 }
2440
2441 ArrayType materializedType = storageSplitType;
2442 Value splitArray;
2443 if (storageSplitType != preciseSplitType) {
2444 ArrayInstantiationInfo instantiationInfo;
2445 switch (inferCommonArrayInstantiation(leafValues, instantiationInfo)) {
2446 case CommonArrayInstantiationStatus::conflict:
2447 // TODO: this POD could be promoted to a complete `struct.def` but that's not easy.
2448 op.emitOpError(
2449 "with POD elements having conflicting affine map instantiations cannot be promoted "
2450 "to higher dimensional array"
2451 );
2452 return failure();
2453 case CommonArrayInstantiationStatus::inferred: {
2454 materializedType = preciseSplitType;
2455 SmallVector<ValueRange> mapOperands;
2456 mapOperands.reserve(instantiationInfo.mapOperandStorage.size());
2457 for (const SmallVector<Value> &values : instantiationInfo.mapOperandStorage) {
2458 mapOperands.push_back(values);
2459 }
2460 CreateArrayOp splitArrayOp = rewriter.create<CreateArrayOp>(
2461 op.getLoc(), materializedType, mapOperands, instantiationInfo.numDimsPerMap
2462 );
2463 preserveDiscardableAttrs(op, splitArrayOp);
2464 splitArray = splitArrayOp;
2465 break;
2466 }
2467 case CommonArrayInstantiationStatus::unavailable:
2468 break;
2469 }
2470 }
2471
2472 if (!splitArray) {
2473 splitArray = createWritableArrayValue(rewriter, op.getLoc(), materializedType);
2474 preserveDiscardableAttrs(op, splitArray.getDefiningOp());
2475 }
2476
2477 for (auto [index, leafValue] : llvm::zip_equal(*elementIndices, leafValues)) {
2478 genArrayWrite(rewriter, op.getLoc(), splitArray, index, leafValue);
2479 }
2480 replacements.push_back(
2481 castValueToTypeIfNeeded(rewriter, op.getLoc(), splitArray, preciseSplitType)
2482 );
2483 }
2484 } else {
2485 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2486 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2487 ArrayType preciseSplitType = llvm::cast<ArrayType>(splitType);
2488 replacements.push_back(createSplitPodArrayReplacement(
2489 op, op.getLoc(), arrTy, id, preciseSplitType, rewriter, mapOperands.ranges,
2490 numDimsPerMap
2491 ));
2492 }
2493 }
2494
2495 if (needsPodArrayShapeCarrier(arrTy)) {
2496 Value shapeCarrier =
2497 materializeArrayLengthCarrier(op.getResult(), arrTy, op.getLoc(), rewriter);
2498 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
2499 replacements.push_back(shapeCarrier);
2500 }
2501
2502 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2503 return success();
2504 }
2505};
2506
2508class SplitPodArrayReadArrayOp : public OpConversionPattern<ReadArrayOp> {
2509public:
2510 using OpConversionPattern<ReadArrayOp>::OpConversionPattern;
2511
2512 static bool legal(ReadArrayOp op) {
2513 return !splittablePodArray(op.getArrRefType()) || shouldDeferPodArrayReadToStep3(op);
2514 }
2515
2516 LogicalResult matchAndRewrite(
2517 ReadArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2518 ) const override {
2519 if (legal(op)) {
2520 return failure();
2521 }
2522 ArrayType arrTy = op.getArrRefType();
2523 PodType podTy = llvm::cast<PodType>(arrTy.getElementType());
2524 SmallVector<RecordChain> splitIds;
2525 SmallVector<Type> splitTypes;
2526 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2527 if (splitTypes.empty()) {
2528 preserveDiscardableAttrs(op, rewriter.replaceOpWithNewOp<NewPodOp>(op, podTy));
2529 return success();
2530 }
2531
2532 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
2533 NewPodOp pod = rewriter.create<NewPodOp>(op.getLoc(), podTy);
2534 preserveDiscardableAttrs(op, pod);
2535 VirtualPodLeafMap leafValues;
2536 auto splitArrRefs = adaptor.getArrRef().take_front(splitIds.size());
2537 for (auto [id, splitType, splitArrRange] :
2538 llvm::zip_equal(splitIds, splitTypes, splitArrRefs)) {
2539 Value leafValue = ArrayAccessOpInterface::genRead(
2540 rewriter, op.getLoc(), getSingleConvertedValue(splitArrRange), indices
2541 );
2542 preserveDiscardableAttrs(op, leafValue.getDefiningOp());
2543 leafValues[id] = tagRaggedNestedLeafValue(
2544 rewriter, op.getLoc(), leafValue, getRaggedNestedLeafAttrName(arrTy, splitType)
2545 );
2546 }
2547
2548 SmallVector<StringAttr> recordChain;
2549 for (RecordAttr record : podTy.getRecords()) {
2550 recordChain.push_back(record.getName());
2551 Value recordValue = rebuildFlattenedPodRecord(
2552 rewriter, op.getLoc(), record.getType(), recordChain, leafValues
2553 );
2554 genWrite(rewriter, op.getLoc(), pod, record.getName(), recordValue);
2555 recordChain.pop_back();
2556 }
2557 rewriter.replaceOp(op, pod);
2558 return success();
2559 }
2560};
2561
2563class SplitPodArrayWriteArrayOp : public OpConversionPattern<WriteArrayOp> {
2564public:
2565 using OpConversionPattern<WriteArrayOp>::OpConversionPattern;
2566
2567 static bool legal(WriteArrayOp op) { return !splittablePodArray(op.getArrRefType()); }
2568
2569 LogicalResult matchAndRewrite(
2570 WriteArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2571 ) const override {
2572 if (legal(op)) {
2573 return failure();
2574 }
2575 ArrayType arrTy = op.getArrRefType();
2576 SmallVector<RecordChain> splitIds;
2577 SmallVector<Type> splitTypes;
2578 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2579 if (splitTypes.empty()) {
2580 rewriter.eraseOp(op);
2581 return success();
2582 }
2583
2584 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
2585 Value podValue = getSingleConvertedValue(adaptor.getRvalue());
2586 auto splitArrRefs = adaptor.getArrRef().take_front(splitIds.size());
2587 for (auto [id, splitArrRange, splitType] :
2588 llvm::zip_equal(splitIds, splitArrRefs, splitTypes)) {
2589 Value leafValue = genReadAlongPath(rewriter, op.getLoc(), podValue, id);
2590 preserveDiscardableAttrs(
2591 op, genArrayWrite(
2592 rewriter, op.getLoc(), getSingleConvertedValue(splitArrRange), indices, leafValue
2593 )
2594 );
2595 }
2596 rewriter.eraseOp(op);
2597 return success();
2598 }
2599};
2600
2602class SplitPodArrayInFuncDefOp : public OpConversionPattern<FuncDefOp> {
2603public:
2604 using OpConversionPattern<FuncDefOp>::OpConversionPattern;
2605
2606 static bool legal(FuncDefOp op) {
2607 return !containsSplittablePodArrayType(op.getArgumentTypes()) &&
2608 !containsSplittablePodArrayType(op.getResultTypes());
2609 }
2610
2611 LogicalResult
2612 matchAndRewrite(FuncDefOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
2613 if (legal(op)) {
2614 return failure();
2615 }
2616 const auto *tyConv = getTypeConverter();
2617 assert(tyConv && "expected pod-array type converter");
2618
2619 FunctionType oldTy = op.getFunctionType();
2620 TypeConverter::SignatureConversion inputConversion(oldTy.getNumInputs());
2621 if (failed(tyConv->convertSignatureArgs(oldTy.getInputs(), inputConversion))) {
2622 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod inputs");
2623 }
2624
2625 SmallVector<Type> newResults;
2626 if (failed(tyConv->convertTypes(oldTy.getResults(), newResults))) {
2627 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod results");
2628 }
2629
2630 if (!op.getBody().empty() &&
2631 failed(rewriter.convertRegionTypes(&op.getBody(), *tyConv, &inputConversion))) {
2632 return rewriter.notifyMatchFailure(op, "failed to convert function body block arguments");
2633 }
2634
2635 SmallVector<size_t> originalInputIdxToSize, originalResultIdxToSize;
2636 SmallVector<Type> newInputs = convertPodArrayTypes(oldTy.getInputs(), &originalInputIdxToSize);
2637 SmallVector<Type> newResultsWithSizeInfo =
2638 convertPodArrayTypes(oldTy.getResults(), &originalResultIdxToSize);
2639 assert(
2640 newResultsWithSizeInfo == newResults &&
2641 "expected array-of-pod type conversion to match function result attr replication"
2642 );
2643 SplitFunctionNameInfo inputNameInfo =
2644 collectSplitFunctionNameInfo(op.getArgumentTypes(), [&](unsigned i) {
2645 return op.getArgNameAttr(i);
2646 }, getSplitPodArrayRecordNameSuffixes);
2647 ArrayAttr resultAttrs = op.getAllResultAttrs();
2648 SplitFunctionNameInfo resultNameInfo =
2649 collectSplitFunctionNameInfo(op.getResultTypes(), [resultAttrs](unsigned i) {
2650 return getAttrAtIndexWithName(resultAttrs, i, RES_NAME_ATTR_NAME);
2651 }, getSplitPodArrayRecordNameSuffixes);
2652
2653 rewriter.modifyOpInPlace(op, [&]() {
2654 op.setFunctionType(FunctionType::get(op.getContext(), newInputs, newResults));
2655 if (ArrayAttr newArgAttrs = replicateFunctionNameAttrsAsNeeded(
2656 op.getArgAttrsAttr(), originalInputIdxToSize, newInputs, ARG_NAME_ATTR_NAME,
2657 inputNameInfo.originalNames, inputNameInfo.existingNames,
2658 inputNameInfo.splitNameSuffixes
2659 )) {
2660 op.setArgAttrsAttr(newArgAttrs);
2661 }
2662 if (ArrayAttr newResAttrs = replicateFunctionNameAttrsAsNeeded(
2663 op.getResAttrsAttr(), originalResultIdxToSize, newResults, RES_NAME_ATTR_NAME,
2664 resultNameInfo.originalNames, resultNameInfo.existingNames,
2665 resultNameInfo.splitNameSuffixes
2666 )) {
2667 op.setResAttrsAttr(newResAttrs);
2668 }
2669 });
2670 return success();
2671 }
2672};
2673
2680static void collectSplitPodArrayOperandValues(
2681 Location loc, Value originalOperand, ValueRange convertedValues,
2682 SmallVectorImpl<Value> &newOperands, ConversionPatternRewriter &rewriter
2683) {
2684 ArrayType arrTy = splittablePodArray(originalOperand.getType());
2685 if (!arrTy) {
2686 llvm::append_range(newOperands, convertedValues);
2687 return;
2688 }
2689
2690 SmallVector<RecordChain> splitIds;
2691 SmallVector<Type> splitTypes;
2692 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
2693 if (splitTypes.empty()) {
2694 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
2695 newOperands.push_back(
2696 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2697 );
2698 return;
2699 }
2700 if (!convertedValues.empty()) {
2701 newOperands.push_back(castValueToTypeIfNeeded(
2702 rewriter, loc, getSingleConvertedValue(convertedValues),
2703 getPodArrayShapeCarrierType(arrTy)
2704 ));
2705 return;
2706 }
2707 newOperands.push_back(materializeArrayLengthCarrier(originalOperand, arrTy, loc, rewriter));
2708 return;
2709 }
2710
2711 ValueRange leafConvertedValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2712 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues);
2713
2714 auto isDirectAggregateToSplitCast = [&leafConvertedValues, &splitTypes]() {
2715 if (leafConvertedValues.empty()) {
2716 return false;
2717 }
2718 auto castOp = leafConvertedValues.front().getDefiningOp<UnrealizedConversionCastOp>();
2719 if (!castOp || castOp->getNumOperands() != 1) {
2720 return false;
2721 }
2722 ArrayType castArrTy = splittablePodArray(castOp.getOperand(0).getType());
2723 size_t expectedResults =
2724 castArrTy ? splitTypes.size() + (needsPodArrayShapeCarrier(castArrTy) ? 1 : 0) : 0;
2725 if (!castArrTy || castOp->getNumResults() != expectedResults) {
2726 return false;
2727 }
2728
2729 return llvm::all_of(llvm::zip_equal(leafConvertedValues, splitTypes), [&castOp](auto pair) {
2730 Value convertedValue = std::get<0>(pair);
2731 Type splitType = std::get<1>(pair);
2732 return convertedValue.getDefiningOp<UnrealizedConversionCastOp>() == castOp &&
2733 typesUnify(convertedValue.getType(), splitType);
2734 });
2735 };
2736 bool directAggregateToSplitCast = isDirectAggregateToSplitCast();
2737
2738 if (!directAggregateToSplitCast && allValueTypesUnifyWithTypes(leafConvertedValues, splitTypes)) {
2739 llvm::append_range(newOperands, leafConvertedValues);
2740 if (carrier) {
2741 newOperands.push_back(
2742 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2743 );
2744 } else if (needsPodArrayShapeCarrier(arrTy)) {
2745 newOperands.push_back(materializeArrayLengthCarrier(originalOperand, arrTy, loc, rewriter));
2746 }
2747 return;
2748 }
2749
2750 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
2751 Value splitValue = genReadAlongPath(rewriter, loc, originalOperand, id);
2752 newOperands.push_back(castValueToTypeIfNeeded(rewriter, loc, splitValue, splitType));
2753 }
2754 if (carrier && !directAggregateToSplitCast) {
2755 newOperands.push_back(
2756 castValueToTypeIfNeeded(rewriter, loc, carrier, getPodArrayShapeCarrierType(arrTy))
2757 );
2758 } else if (needsPodArrayShapeCarrier(arrTy)) {
2759 RecordChain carrierId({getPodArrayShapeCarrierMarker(rewriter.getContext())}, true);
2760 Value splitCarrier = genReadAlongPath(rewriter, loc, originalOperand, carrierId);
2761 newOperands.push_back(
2762 castValueToTypeIfNeeded(rewriter, loc, splitCarrier, getPodArrayShapeCarrierType(arrTy))
2763 );
2764 }
2765}
2766
2768class SplitPodArrayInUnifiableCastOp : public OpConversionPattern<UnifiableCastOp> {
2769public:
2770 using OpConversionPattern<UnifiableCastOp>::OpConversionPattern;
2771
2772 static bool legal(UnifiableCastOp op) {
2773 return !splittablePodArray(op.getType()) && !splittablePodArray(op.getInput().getType());
2774 }
2775
2776 LogicalResult matchAndRewrite(
2777 UnifiableCastOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2778 ) const override {
2779 if (legal(op)) {
2780 return failure();
2781 }
2782
2783 ArrayType inputArrTy = splittablePodArray(op.getInput().getType());
2784 ArrayType resultArrTy = splittablePodArray(op.getType());
2785
2786 // Lowering a split array-of-POD input to one non-array SSA value would require either
2787 // materializing the original aggregate or rewriting the surrounding function/template
2788 // signature to thread separate generic pieces.
2789 if (inputArrTy && !resultArrTy) {
2790 return rewriter.notifyMatchFailure(
2791 op, "array-of-pod input to non-array result requires aggregate materialization or "
2792 "signature/template rewriting"
2793 );
2794 }
2795
2796 if (!inputArrTy) {
2797 return rewriter.notifyMatchFailure(
2798 op, "generic input to array-of-pod result requires signature/template rewriting"
2799 );
2800 }
2801
2802 SmallVector<RecordChain> inputSplitIds;
2803 SmallVector<Type> inputSplitTypes;
2804 splitPodArrayTypeTo(inputArrTy, inputSplitTypes, &inputSplitIds);
2805
2806 SmallVector<RecordChain> resultSplitIds;
2807 SmallVector<Type> resultSplitTypes;
2808 splitPodArrayTypeTo(resultArrTy, resultSplitTypes, &resultSplitIds);
2809
2810 if (inputSplitIds != resultSplitIds) {
2811 return rewriter.notifyMatchFailure(
2812 op, "array-of-pod cast changed POD leaf structure unexpectedly"
2813 );
2814 }
2815
2816 SmallVector<Value> splitInputs;
2817 collectSplitPodArrayOperandValues(
2818 op.getLoc(), op.getInput(), adaptor.getInput(), splitInputs, rewriter
2819 );
2820 ValueRange splitInputLeaves = getConvertedPodArrayLeafValues(inputArrTy, splitInputs);
2821 if (resultSplitTypes.empty()) {
2822 if (splitInputs.size() != 1) {
2823 return rewriter.notifyMatchFailure(
2824 op, "expected one shape carrier for zero-leaf array-of-pod cast"
2825 );
2826 }
2827 Value replacement = castValueToTypeIfNeeded(
2828 rewriter, op.getLoc(), splitInputs.front(), getPodArrayShapeCarrierType(resultArrTy)
2829 );
2830 if (replacement != splitInputs.front()) {
2831 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2832 }
2833 rewriter.replaceOp(op, replacement);
2834 return success();
2835 }
2836 if (splitInputLeaves.size() != resultSplitTypes.size()) {
2837 return rewriter.notifyMatchFailure(
2838 op, "failed to collect one split input per array-of-pod cast leaf"
2839 );
2840 }
2841
2842 SmallVector<Value> replacements;
2843 replacements.reserve(
2844 resultSplitTypes.size() + (needsPodArrayShapeCarrier(resultArrTy) ? 1 : 0)
2845 );
2846 for (auto [splitInput, resultSplitType] : llvm::zip_equal(splitInputLeaves, resultSplitTypes)) {
2847 Value replacement =
2848 castValueToTypeIfNeeded(rewriter, op.getLoc(), splitInput, resultSplitType);
2849 if (replacement != splitInput) {
2850 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2851 }
2852 replacements.push_back(replacement);
2853 }
2854 if (needsPodArrayShapeCarrier(resultArrTy)) {
2855 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(inputArrTy, splitInputs);
2856 if (!carrier) {
2857 carrier = materializeArrayLengthCarrier(op.getInput(), inputArrTy, op.getLoc(), rewriter);
2858 }
2859 Value replacement = castValueToTypeIfNeeded(
2860 rewriter, op.getLoc(), carrier, getPodArrayShapeCarrierType(resultArrTy)
2861 );
2862 if (replacement != carrier) {
2863 preserveDiscardableAttrs(op, replacement.getDefiningOp());
2864 }
2865 replacements.push_back(replacement);
2866 }
2867
2868 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
2869 return success();
2870 }
2871};
2872
2874class SplitPodArrayInReturnOp : public OpConversionPattern<ReturnOp> {
2875public:
2876 using OpConversionPattern<ReturnOp>::OpConversionPattern;
2877
2878 static bool legal(ReturnOp op) {
2879 return !containsSplittablePodArrayType(op.getOperands().getTypes());
2880 }
2881
2882 LogicalResult matchAndRewrite(
2883 ReturnOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2884 ) const override {
2885 if (legal(op)) {
2886 return failure();
2887 }
2888 SmallVector<Value> newOperands;
2889 for (auto [operand, convertedValues] :
2890 llvm::zip_equal(op.getOperands(), adaptor.getOperands())) {
2891 collectSplitPodArrayOperandValues(
2892 op.getLoc(), operand, convertedValues, newOperands, rewriter
2893 );
2894 }
2895 preserveDiscardableAttrs(
2896 op, rewriter.replaceOpWithNewOp<ReturnOp>(op, ValueRange(newOperands))
2897 );
2898 return success();
2899 }
2900};
2901
2903class SplitPodArrayInCallOp : public OpConversionPattern<CallOp> {
2904public:
2905 using OpConversionPattern<CallOp>::OpConversionPattern;
2906
2907 static bool legal(CallOp op) {
2908 return !containsSplittablePodArrayType(op.getArgOperands().getTypes()) &&
2909 !containsSplittablePodArrayType(op.getResultTypes());
2910 }
2911
2912 LogicalResult matchAndRewrite(
2913 CallOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
2914 ) const override {
2915 if (legal(op)) {
2916 return failure();
2917 }
2918 const auto *tyConv = getTypeConverter();
2919 assert(tyConv && "expected pod-array type converter");
2920
2921 SmallVector<Type> newResultTypes;
2922 if (failed(tyConv->convertTypes(op.getResultTypes(), newResultTypes))) {
2923 return rewriter.notifyMatchFailure(op, "failed to convert array-of-pod call results");
2924 }
2925
2926 FlattenedConvertedValueRangeStorage mapOperands(adaptor.getMapOperands());
2927
2928 SmallVector<Value> newArgOperands;
2929 for (auto [operand, convertedValues] :
2930 llvm::zip_equal(op.getArgOperands(), adaptor.getArgOperands())) {
2931 collectSplitPodArrayOperandValues(
2932 op.getLoc(), operand, convertedValues, newArgOperands, rewriter
2933 );
2934 }
2936 op.getLoc(), newResultTypes, op, mapOperands.ranges, newArgOperands, rewriter
2937 );
2938
2939 SmallVector<SmallVector<Value>> replacementStorage;
2940 replacementStorage.reserve(op.getNumResults());
2941 auto newResultIt = newCall.getResults().begin();
2942 for (Type oldResultType : op.getResultTypes()) {
2943 SmallVector<Type> convertedTypes;
2944 (void)convertPodArrayTypeTo(oldResultType, convertedTypes);
2945 SmallVector<Value> replacementsForResult;
2946 replacementsForResult.reserve(convertedTypes.size());
2947 for (size_t i = 0; i < convertedTypes.size(); ++i) {
2948 replacementsForResult.push_back(*newResultIt);
2949 ++newResultIt;
2950 }
2951 replacementStorage.push_back(std::move(replacementsForResult));
2952 }
2953
2954 SmallVector<ValueRange> replacements;
2955 replacements.reserve(replacementStorage.size());
2956 for (const SmallVector<Value> &values : replacementStorage) {
2957 replacements.push_back(values);
2958 }
2959 rewriter.replaceOpWithMultiple(op, replacements);
2960 return success();
2961 }
2962};
2963
2965static Value getPodArrayEqualityShapeSource(
2966 ArrayType arrTy, Value originalValue, ValueRange convertedValues, Location loc,
2967 OpBuilder &rewriter
2968) {
2969 if (Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, convertedValues)) {
2970 return carrier;
2971 }
2972
2973 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2974 if (!leafValues.empty()) {
2975 return leafValues.front();
2976 }
2977
2978 if (!convertedValues.empty()) {
2979 return getSingleConvertedValue(convertedValues);
2980 }
2981
2982 return materializeArrayLengthCarrier(originalValue, arrTy, loc, rewriter);
2983}
2984
2986static SmallVector<Value> collectCompatiblePodArrayEqualityLeaves(
2987 Location loc, Value originalValue, ValueRange convertedValues, ArrayType arrTy,
2988 ArrayType peerArrTy, ConversionPatternRewriter &rewriter,
2989 CompatiblePodLeafMaterializationMap &materializedLeaves
2990) {
2991 if (arrTy) {
2992 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
2993 size_t leafCount = getSplitPodArrayLeafCount(arrTy);
2994 if (leafValues.size() == leafCount || leafCount == 0) {
2995 return SmallVector<Value>(leafValues.begin(), leafValues.end());
2996 }
2997
2998 return materializeCompatiblePodArrayLeafValues(loc, originalValue, arrTy, rewriter);
2999 }
3000
3001 if (!peerArrTy || convertedValues.size() != 1) {
3002 if (!peerArrTy || !convertedValues.empty()) {
3003 return SmallVector<Value>(convertedValues.begin(), convertedValues.end());
3004 }
3005 }
3006
3007 Value source = convertedValues.empty() ? originalValue : getSingleConvertedValue(convertedValues);
3008 SmallVector<Value> materialized;
3009 llvm::append_range(
3010 materialized, getOrMaterializeCompatiblePodArrayLeafValues(
3011 loc, source, peerArrTy, rewriter, materializedLeaves
3012 )
3013 );
3014 return materialized;
3015}
3016
3018static ArrayType getCompatiblePodArrayType(ArrayType arrTy, PodType elemPodTy) {
3019 if (ArrayType concreteArrTy = splittablePodArray(arrTy)) {
3020 return concreteArrTy;
3021 }
3022 return ArrayType::get(elemPodTy, arrTy.getDimensionSizes());
3023}
3024
3026class SplitPodArrayInEmitEqualityOp : public OpConversionPattern<constrain::EmitEqualityOp> {
3027 CompatiblePodLeafMaterializationMap &materializedLeaves;
3028
3029public:
3030 SplitPodArrayInEmitEqualityOp(
3031 TypeConverter &tyConv, MLIRContext *ctx,
3032 CompatiblePodLeafMaterializationMap &materializedLeafMap
3033 )
3034 : OpConversionPattern<constrain::EmitEqualityOp>(tyConv, ctx),
3035 materializedLeaves(materializedLeafMap) {}
3036
3037 static bool legal(constrain::EmitEqualityOp op) {
3038 return !containsSplittablePodArrayType(op->getOperandTypes());
3039 }
3040
3041 LogicalResult matchAndRewrite(
3042 constrain::EmitEqualityOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3043 ) const override {
3044 if (legal(op)) {
3045 return failure();
3046 }
3047
3048 ArrayType lhsTy = splittablePodArray(op.getLhs().getType());
3049 ArrayType rhsTy = splittablePodArray(op.getRhs().getType());
3050 StringRef raggedKind = lhsTy ? getRaggedNestedLeafKind(lhsTy) : StringRef {};
3051 if (raggedKind.empty() && rhsTy) {
3052 raggedKind = getRaggedNestedLeafKind(rhsTy);
3053 }
3054 if (!raggedKind.empty()) {
3055 return op.emitOpError()
3056 << "cannot lower nested " << raggedKind
3057 << " array leaf equality for array-of-POD without per-element shape witnesses";
3058 }
3059
3060 bool lhsNeedsShapeCheck = lhsTy && needsPodArrayShapeCarrier(lhsTy);
3061 bool rhsNeedsShapeCheck = rhsTy && needsPodArrayShapeCarrier(rhsTy);
3062 SmallVector<Value> lhsCompatibleConvertedValues;
3063 SmallVector<Value> rhsCompatibleConvertedValues;
3064 if (rhsTy && rhsNeedsShapeCheck && !lhsTy &&
3065 (adaptor.getLhs().empty() || adaptor.getLhs().size() == 1)) {
3066 Value lhsSource =
3067 adaptor.getLhs().empty() ? op.getLhs() : getSingleConvertedValue(adaptor.getLhs());
3068 lhsCompatibleConvertedValues =
3069 materializeCompatiblePodArrayConvertedValues(op.getLoc(), lhsSource, rhsTy, rewriter);
3070 }
3071 if (lhsTy && lhsNeedsShapeCheck && !rhsTy &&
3072 (adaptor.getRhs().empty() || adaptor.getRhs().size() == 1)) {
3073 Value rhsSource =
3074 adaptor.getRhs().empty() ? op.getRhs() : getSingleConvertedValue(adaptor.getRhs());
3075 rhsCompatibleConvertedValues =
3076 materializeCompatiblePodArrayConvertedValues(op.getLoc(), rhsSource, lhsTy, rewriter);
3077 }
3078
3079 ArrayType shapeCheckTy = lhsTy ? lhsTy : rhsTy;
3080 bool needsShapeCheck = lhsTy ? lhsNeedsShapeCheck : rhsNeedsShapeCheck;
3081 if (lhsTy && rhsTy) {
3082 if (lhsTy.getDimensionSizes().size() != rhsTy.getDimensionSizes().size()) {
3083 return rewriter.notifyMatchFailure(
3084 op, "expected array-of-pod equality operands with matching rank"
3085 );
3086 }
3087 needsShapeCheck = lhsNeedsShapeCheck || rhsNeedsShapeCheck;
3088 }
3089
3090 if (shapeCheckTy && needsShapeCheck) {
3091 Value lhsShapeSource = getPodArrayEqualityShapeSource(
3092 lhsTy ? lhsTy : rhsTy, op.getLhs(),
3093 lhsTy ? adaptor.getLhs() : ValueRange(lhsCompatibleConvertedValues), op.getLoc(), rewriter
3094 );
3095 Value rhsShapeSource = getPodArrayEqualityShapeSource(
3096 rhsTy ? rhsTy : lhsTy, op.getRhs(),
3097 rhsTy ? adaptor.getRhs() : ValueRange(rhsCompatibleConvertedValues), op.getLoc(), rewriter
3098 );
3099
3100 for (size_t dim = 0, rank = shapeCheckTy.getDimensionSizes().size(); dim < rank; ++dim) {
3101 Value dimVal = rewriter.create<arith::ConstantOp>(
3102 op.getLoc(), rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3103 );
3104 Value lhsLen = rewriter.create<ArrayLengthOp>(op.getLoc(), lhsShapeSource, dimVal);
3105 Value rhsLen = rewriter.create<ArrayLengthOp>(op.getLoc(), rhsShapeSource, dimVal);
3106 preserveDiscardableAttrs(
3107 op, rewriter.create<constrain::EmitEqualityOp>(op.getLoc(), lhsLen, rhsLen)
3108 );
3109 }
3110 }
3111
3112 SmallVector<Value> lhsLeaves;
3113 if (!lhsTy && rhsTy && !lhsCompatibleConvertedValues.empty()) {
3114 ValueRange materializedLeavesForLhs =
3115 getConvertedPodArrayLeafValues(rhsTy, lhsCompatibleConvertedValues);
3116 llvm::append_range(lhsLeaves, materializedLeavesForLhs);
3117 } else {
3118 lhsLeaves = collectCompatiblePodArrayEqualityLeaves(
3119 op.getLoc(), op.getLhs(), adaptor.getLhs(), lhsTy, rhsTy, rewriter, materializedLeaves
3120 );
3121 }
3122
3123 SmallVector<Value> rhsLeaves;
3124 if (!rhsTy && lhsTy && !rhsCompatibleConvertedValues.empty()) {
3125 ValueRange materializedLeavesForRhs =
3126 getConvertedPodArrayLeafValues(lhsTy, rhsCompatibleConvertedValues);
3127 llvm::append_range(rhsLeaves, materializedLeavesForRhs);
3128 } else {
3129 rhsLeaves = collectCompatiblePodArrayEqualityLeaves(
3130 op.getLoc(), op.getRhs(), adaptor.getRhs(), rhsTy, lhsTy, rewriter, materializedLeaves
3131 );
3132 }
3133 if (lhsLeaves.size() != rhsLeaves.size()) {
3134 return rewriter.notifyMatchFailure(
3135 op, "expected array-of-pod equality operands to expand to the same number of leaves"
3136 );
3137 }
3138
3139 for (auto [lhs, rhs] : llvm::zip_equal(lhsLeaves, rhsLeaves)) {
3140 preserveDiscardableAttrs(
3141 op, rewriter.create<constrain::EmitEqualityOp>(op.getLoc(), lhs, rhs)
3142 );
3143 }
3144 rewriter.eraseOp(op);
3145 return success();
3146 }
3147};
3148
3171class SplitPodArrayInEmitContainmentOp : public OpConversionPattern<constrain::EmitContainmentOp> {
3172 CompatiblePodLeafMaterializationMap &materializedLeaves;
3173
3174public:
3175 SplitPodArrayInEmitContainmentOp(
3176 TypeConverter &tyConv, MLIRContext *ctx,
3177 CompatiblePodLeafMaterializationMap &materializedLeafMap
3178 )
3179 : OpConversionPattern<constrain::EmitContainmentOp>(tyConv, ctx),
3180 materializedLeaves(materializedLeafMap) {}
3181
3182 static bool legal(constrain::EmitContainmentOp op) {
3183 return !containsSplittablePodArrayType(op->getOperandTypes()) &&
3184 getTaggedRaggedNestedLeafKind(op.getLhs()).empty() &&
3185 getTaggedRaggedNestedLeafKind(op.getRhs()).empty();
3186 }
3187
3189 static SmallVector<Value> collectContainmentLeaves(
3190 Location loc, Value originalOperand, ValueRange convertedValues,
3191 ConversionPatternRewriter &rewriter,
3192 CompatiblePodLeafMaterializationMap *materializedLeaves = nullptr,
3193 std::optional<PodType> compatiblePodTy = std::nullopt
3194 ) {
3195 if (ArrayType arrTy = splittablePodArray(originalOperand.getType())) {
3196 ValueRange leafValues = getConvertedPodArrayLeafValues(arrTy, convertedValues);
3197 if (hasZeroLeafPodArraySplit(arrTy)) {
3198 return {};
3199 }
3200 size_t leafCount = getSplitPodArrayLeafCount(arrTy);
3201 if (leafValues.size() == leafCount) {
3202 return SmallVector<Value>(leafValues.begin(), leafValues.end());
3203 }
3204 return materializeCompatiblePodArrayLeafValues(loc, originalOperand, arrTy, rewriter);
3205 }
3206
3207 if (splittablePod(originalOperand.getType())) {
3208 SmallVector<Value> podLeaves;
3209 processInputOperand(loc, getSingleConvertedValue(convertedValues), podLeaves, rewriter);
3210 return podLeaves;
3211 }
3212
3213 if (materializedLeaves && compatiblePodTy &&
3214 (convertedValues.empty() || convertedValues.size() == 1)) {
3215 Value source =
3216 convertedValues.empty() ? originalOperand : getSingleConvertedValue(convertedValues);
3217 ArrayRef<Value> leaves = getOrMaterializeCompatiblePodLeafValues(
3218 loc, source, *compatiblePodTy, rewriter, *materializedLeaves
3219 );
3220 return SmallVector<Value>(leaves.begin(), leaves.end());
3221 }
3222
3223 return SmallVector<Value>(convertedValues.begin(), convertedValues.end());
3224 }
3225
3226 LogicalResult matchAndRewrite(
3227 constrain::EmitContainmentOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3228 ) const override {
3229 if (failed(rejectRaggedNestedLeafContainment(op))) {
3230 return failure();
3231 }
3232
3233 if (legal(op)) {
3234 return failure();
3235 }
3236
3237 Location loc = op.getLoc();
3238 ArrayType lhsTy = op.getLhs().getType();
3239 Type rhsTy = op.getRhs().getType();
3240
3241 size_t lhsRank = lhsTy.getDimensionSizes().size();
3242 size_t rhsRank = 0;
3243 if (auto rhsArrTy = llvm::dyn_cast<ArrayType>(rhsTy)) {
3244 rhsRank = rhsArrTy.getDimensionSizes().size();
3245 }
3246 assert(lhsRank >= rhsRank && "constrain.in verifier should reject higher-rank rhs arrays");
3247 size_t selectedDims = lhsRank - rhsRank;
3248
3249 ArrayType lhsPodArrTy = splittablePodArray(lhsTy);
3250 ArrayType rhsPodArrTy = splittablePodArray(rhsTy);
3251 assert(
3252 (lhsPodArrTy || rhsPodArrTy) &&
3253 "containment rewrite requires at least one concrete array-of-POD operand"
3254 );
3255
3256 PodType compatibleElemPodTy = lhsPodArrTy ? llvm::cast<PodType>(lhsPodArrTy.getElementType())
3257 : llvm::cast<PodType>(rhsPodArrTy.getElementType());
3258 ArrayType compatibleLhsTy = getCompatiblePodArrayType(lhsTy, compatibleElemPodTy);
3259 ArrayType rhsArrTy = llvm::dyn_cast<ArrayType>(rhsTy);
3260 ArrayType compatibleRhsArrTy =
3261 rhsArrTy ? getCompatiblePodArrayType(rhsArrTy, compatibleElemPodTy) : ArrayType();
3262
3263 SmallVector<Value> lhsCompatibleConvertedValues;
3264 if (!lhsPodArrTy) {
3265 if (!adaptor.getLhs().empty() && adaptor.getLhs().size() != 1) {
3266 return rewriter.notifyMatchFailure(
3267 op, "expected a single converted lhs value to materialize generic containment leaves"
3268 );
3269 }
3270 Value lhsSource =
3271 adaptor.getLhs().empty() ? op.getLhs() : getSingleConvertedValue(adaptor.getLhs());
3272 lhsCompatibleConvertedValues =
3273 materializeCompatiblePodArrayConvertedValues(loc, lhsSource, compatibleLhsTy, rewriter);
3274 }
3275
3276 SmallVector<Value> rhsCompatibleConvertedValues;
3277 if (compatibleRhsArrTy && !rhsPodArrTy) {
3278 if (!adaptor.getRhs().empty() && adaptor.getRhs().size() != 1) {
3279 return rewriter.notifyMatchFailure(
3280 op, "expected a single converted rhs value to materialize generic containment leaves"
3281 );
3282 }
3283 Value rhsSource =
3284 adaptor.getRhs().empty() ? op.getRhs() : getSingleConvertedValue(adaptor.getRhs());
3285 rhsCompatibleConvertedValues = materializeCompatiblePodArrayConvertedValues(
3286 loc, rhsSource, compatibleRhsArrTy, rewriter
3287 );
3288 }
3289
3290 SmallVector<Value> lhsLeaves;
3291 if (!lhsCompatibleConvertedValues.empty()) {
3292 llvm::append_range(
3293 lhsLeaves, getConvertedPodArrayLeafValues(compatibleLhsTy, lhsCompatibleConvertedValues)
3294 );
3295 } else {
3296 lhsLeaves = collectContainmentLeaves(loc, op.getLhs(), adaptor.getLhs(), rewriter);
3297 }
3298
3299 SmallVector<Value> rhsLeaves;
3300 if (!rhsCompatibleConvertedValues.empty()) {
3301 llvm::append_range(
3302 rhsLeaves,
3303 getConvertedPodArrayLeafValues(compatibleRhsArrTy, rhsCompatibleConvertedValues)
3304 );
3305 } else {
3306 rhsLeaves = collectContainmentLeaves(
3307 loc, op.getRhs(), adaptor.getRhs(), rewriter, &materializedLeaves, compatibleElemPodTy
3308 );
3309 }
3310 if (lhsLeaves.size() != rhsLeaves.size()) {
3311 return rewriter.notifyMatchFailure(
3312 op, "expected array-of-pod containment operands to expand to the same number of leaves"
3313 );
3314 }
3315
3316 auto getShapeSource =
3317 [&loc, &rewriter](ArrayType arrTy, Value originalValue, ValueRange convertedValues) {
3318 if (Value carrier = getConvertedPodArrayShapeSource(arrTy, convertedValues)) {
3319 return carrier;
3320 }
3321
3322 if (!convertedValues.empty()) {
3323 return getSingleConvertedValue(convertedValues);
3324 }
3325
3326 return materializeArrayLengthCarrier(originalValue, arrTy, loc, rewriter);
3327 };
3328
3329 Value shapeCarrier = getShapeSource(
3330 compatibleLhsTy, op.getLhs(),
3331 lhsCompatibleConvertedValues.empty() ? adaptor.getLhs()
3332 : ValueRange(lhsCompatibleConvertedValues)
3333 );
3334 Value zero = rewriter.create<arith::ConstantOp>(loc, rewriter.getIndexAttr(0));
3335 Value trueVal = rewriter.create<arith::ConstantOp>(
3336 loc, IntegerAttr::get(IntegerType::get(rewriter.getContext(), 1), 1)
3337 );
3338
3339 SmallVector<Value> selectedIndices;
3340 selectedIndices.reserve(selectedDims);
3341 for (size_t dim = 0; dim < selectedDims; ++dim) {
3342 Value idx = rewriter.create<NonDetOp>(loc, IndexType::get(rewriter.getContext()));
3343 Value dimVal = rewriter.create<arith::ConstantOp>(
3344 loc, rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3345 );
3346 Value dimLen = rewriter.create<ArrayLengthOp>(loc, shapeCarrier, dimVal);
3347
3348 Value nonNegative = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::sge, idx, zero);
3349 rewriter.create<constrain::EmitEqualityOp>(loc, nonNegative, trueVal);
3350
3351 Value inRange = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::slt, idx, dimLen);
3352 rewriter.create<constrain::EmitEqualityOp>(loc, inRange, trueVal);
3353
3354 selectedIndices.push_back(idx);
3355 }
3356
3357 bool lhsNeedsShapeCheck = rhsArrTy && needsPodArrayShapeCarrier(compatibleLhsTy);
3358 bool rhsNeedsShapeCheck = compatibleRhsArrTy && needsPodArrayShapeCarrier(compatibleRhsArrTy);
3359 if (compatibleRhsArrTy && (lhsNeedsShapeCheck || rhsNeedsShapeCheck)) {
3360 Value rhsShapeSource = getShapeSource(
3361 compatibleRhsArrTy, op.getRhs(),
3362 rhsCompatibleConvertedValues.empty() ? adaptor.getRhs()
3363 : ValueRange(rhsCompatibleConvertedValues)
3364 );
3365 Value selectedShapeSource =
3366 selectedIndices.empty()
3367 ? shapeCarrier
3368 : ArrayAccessOpInterface::genRead(rewriter, loc, shapeCarrier, selectedIndices);
3369 for (size_t dim = 0; dim < rhsRank; ++dim) {
3370 Value dimVal = rewriter.create<arith::ConstantOp>(
3371 loc, rewriter.getIndexAttr(llzk::checkedCast<int64_t>(dim))
3372 );
3373 Value lhsLen = rewriter.create<ArrayLengthOp>(loc, selectedShapeSource, dimVal);
3374 Value rhsLen = rewriter.create<ArrayLengthOp>(loc, rhsShapeSource, dimVal);
3375 preserveDiscardableAttrs(
3376 op, rewriter.create<constrain::EmitEqualityOp>(loc, lhsLen, rhsLen)
3377 );
3378 }
3379 }
3380
3381 if (lhsLeaves.empty() && rhsLeaves.empty()) {
3382 if (rhsArrTy) {
3383 Value rhsShapeCarrier = getShapeSource(rhsArrTy, op.getRhs(), adaptor.getRhs());
3384 Value selectedShape =
3385 selectedIndices.empty()
3386 ? shapeCarrier
3387 : rewriter.create<ExtractArrayOp>(
3388 loc, getPodArrayShapeCarrierType(rhsArrTy), shapeCarrier, selectedIndices
3389 );
3390 preserveDiscardableAttrs(
3391 op, rewriter.create<constrain::EmitEqualityOp>(loc, selectedShape, rhsShapeCarrier)
3392 );
3393 }
3394 rewriter.eraseOp(op);
3395 return success();
3396 }
3397
3398 for (auto [lhsLeaf, rhsLeaf] : llvm::zip_equal(lhsLeaves, rhsLeaves)) {
3399 Value selectedLhs = lhsLeaf;
3400 if (auto rhsLeafArrTy = llvm::dyn_cast<ArrayType>(rhsLeaf.getType())) {
3401 if (!selectedIndices.empty()) {
3402 selectedLhs =
3403 rewriter.create<ExtractArrayOp>(loc, rhsLeafArrTy, lhsLeaf, selectedIndices);
3404 }
3405 } else {
3406 selectedLhs =
3407 rewriter.create<ReadArrayOp>(loc, rhsLeaf.getType(), lhsLeaf, selectedIndices);
3408 }
3409 preserveDiscardableAttrs(
3410 op, rewriter.create<constrain::EmitEqualityOp>(loc, selectedLhs, rhsLeaf)
3411 );
3412 }
3413
3414 rewriter.eraseOp(op);
3415 return success();
3416 }
3417};
3418
3425static Value selectArrayLengthShapeSource(
3426 ArrayLengthOp op, ValueRange convertedArrRefs, ConversionPatternRewriter &rewriter
3427) {
3428 if (Value v = getConvertedPodArrayShapeSource(op.getArrRefType(), convertedArrRefs)) {
3429 return v;
3430 }
3431
3432 return materializeArrayLengthCarrier(op.getArrRef(), op.getArrRefType(), op.getLoc(), rewriter);
3433}
3434
3435static Value tryResolveReadPodArrayShapeSource(
3436 ReadPodOp readOp, ArrayType arrTy, const VirtualPodValueMap &virtualPods, Location loc,
3437 RewriterBase &rewriter
3438);
3439
3441class RejectRaggedNestedLeafArrayLengthOp : public OpConversionPattern<ArrayLengthOp> {
3442public:
3443 using OpConversionPattern<ArrayLengthOp>::OpConversionPattern;
3444
3445 static bool legal(ArrayLengthOp op) {
3446 return getTaggedRaggedNestedLeafKind(op.getArrRef()).empty();
3447 }
3448
3449 LogicalResult
3450 matchAndRewrite(ArrayLengthOp op, OneToNOpAdaptor, ConversionPatternRewriter &) const override {
3451 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getArrRef());
3452 if (raggedKind.empty()) {
3453 return failure();
3454 }
3455 return op.emitOpError() << "cannot lower nested " << raggedKind
3456 << " array leaf length after reading an array-of-POD element without "
3457 "per-element shape witnesses";
3458 }
3459};
3460
3462class SplitPodArrayLengthOp : public OpConversionPattern<ArrayLengthOp> {
3463public:
3464 using OpConversionPattern<ArrayLengthOp>::OpConversionPattern;
3465
3466 static bool legal(ArrayLengthOp op) {
3467 return RejectRaggedNestedLeafArrayLengthOp::legal(op) &&
3468 (!splittablePodArray(op.getArrRefType()) || op->hasAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR));
3469 }
3470
3471 LogicalResult matchAndRewrite(
3472 ArrayLengthOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3473 ) const override {
3474 if (legal(op)) {
3475 return failure();
3476 }
3477 if (shouldDeferPodArrayLengthToStep3(op)) {
3478 auto deferred = rewriter.create<ArrayLengthOp>(
3479 op.getLoc(), op.getArrRef(), getSingleConvertedValue(adaptor.getDim())
3480 );
3481 preserveDiscardableAttrs(op, deferred);
3482 deferred->setAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR, UnitAttr::get(op.getContext()));
3483 rewriter.replaceOp(op, deferred.getResult());
3484 return success();
3485 }
3486 Value arrRef = selectArrayLengthShapeSource(op, adaptor.getArrRef(), rewriter);
3487 preserveDiscardableAttrs(
3488 op, rewriter.replaceOpWithNewOp<ArrayLengthOp>(
3489 op, arrRef, getSingleConvertedValue(adaptor.getDim())
3490 )
3491 );
3492 return success();
3493 }
3494};
3495
3501struct DeferredPodArrayBacking {
3502 SmallVector<Value> leafArrays;
3503 Value shapeCarrier;
3504};
3505
3506using DeferredPodArrayBackingMap = DenseMap<Value, DeferredPodArrayBacking>;
3507
3509static void setDeferredPodArrayBackingInsertionPoint(ReadPodOp readOp, OpBuilder &bldr) {
3510 if (Operation *loopOp = llzk::pod::detail::findNearestLoopCarriedPodAccess(readOp)) {
3511 bldr.setInsertionPoint(loopOp);
3512 } else {
3513 bldr.setInsertionPointAfter(readOp);
3514 }
3515}
3516
3518static DeferredPodArrayBacking &materializeDeferredPodArrayBacking(
3519 ReadPodOp readOp, ArrayType arrTy, ArrayRef<Type> splitTypes,
3520 DeferredPodArrayBackingMap &deferredPodArrays, Location loc, OpBuilder &bldr,
3521 bool requireLeafArrays, bool requireShapeCarrier
3522) {
3523 auto [it, inserted] = deferredPodArrays.try_emplace(readOp.getResult());
3524 DeferredPodArrayBacking &backing = it->second;
3525
3526 bool needsShapeCarrier = requireShapeCarrier && needsPodArrayShapeCarrier(arrTy);
3527 bool missingLeafArrays = requireLeafArrays && backing.leafArrays.empty();
3528 bool missingShapeCarrier = needsShapeCarrier && !backing.shapeCarrier;
3529 if (missingLeafArrays || missingShapeCarrier) {
3530 OpBuilder::InsertionGuard guard(bldr);
3531 setDeferredPodArrayBackingInsertionPoint(readOp, bldr);
3532
3533 if (missingLeafArrays) {
3534 backing.leafArrays.reserve(splitTypes.size());
3535 for (Type splitType : splitTypes) {
3536 backing.leafArrays.push_back(
3537 createWritableArrayValue(bldr, loc, llvm::cast<ArrayType>(splitType))
3538 );
3539 }
3540 }
3541
3542 if (missingShapeCarrier) {
3543 backing.shapeCarrier = materializeArrayLengthCarrier(readOp.getResult(), arrTy, loc, bldr);
3544 }
3545 }
3546
3547 if (inserted || requireLeafArrays) {
3548 assert(
3549 backing.leafArrays.size() == splitTypes.size() &&
3550 "cached split POD arrays must match the rewritten read arity"
3551 );
3552 }
3553 return backing;
3554}
3555
3557struct Step3Resolver {
3558 VirtualPodValueMap virtualPods;
3559 CompatiblePodLeafMaterializationMap materializedLeaves;
3560 DeferredPodArrayBackingMap deferredPodArrays;
3561
3562 void rehydrateVirtualPodPlaceholders(ModuleOp modOp);
3563 void addPreConversionPatterns(RewritePatternSet &patterns);
3564 void addConversionPatterns(
3565 RewritePatternSet &patterns, SymbolTableCollection &symTables,
3566 const MemberReplacementMap &memberRepMap
3567 );
3568 void addLateResolutionPatterns(RewritePatternSet &patterns);
3569 void addPostConversionPatterns(RewritePatternSet &patterns);
3570 void configureLateVirtualPodLegality(ConversionTarget &target) const;
3571 bool hasResolvableLateVirtualPodOps(ModuleOp modOp) const;
3572 void materializeRemainingVirtualPods(ModuleOp modOp);
3573};
3574
3576class ResolvePodReadBackedArrayLengthOp final : public OpConversionPattern<ArrayLengthOp> {
3577 Step3Resolver &resolver;
3578
3579public:
3580 ResolvePodReadBackedArrayLengthOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
3581 : OpConversionPattern<ArrayLengthOp>(ctx), resolver(step3Resolver) {}
3582
3583 static bool legal(ArrayLengthOp op) { return !op->hasAttr(DEFERRED_POD_ARRAY_LENGTH_ATTR); }
3584
3585 LogicalResult
3586 matchAndRewrite(ArrayLengthOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
3587 ArrayType arrTy = splittablePodArray(op.getArrRefType());
3588 if (!arrTy || !shouldDeferPodArrayLengthToStep3(op)) {
3589 return failure();
3590 }
3591
3592 Value shapeSource;
3593 ReadPodOp readOp = getReadPodBacking(op.getArrRef());
3594 assert(readOp && "deferred POD-backed array.len must still trace to pod.read");
3595 shapeSource = tryResolveReadPodArrayShapeSource(
3596 readOp, arrTy, resolver.virtualPods, op.getLoc(), rewriter
3597 );
3598
3599 if (!shapeSource) {
3600 if (needsPodArrayShapeCarrier(arrTy) && isFreshUnwrittenPodRead(readOp)) {
3601 shapeSource =
3602 materializeDeferredPodArrayBacking(
3603 readOp, arrTy, /*splitTypes=*/ArrayRef<Type> {}, resolver.deferredPodArrays,
3604 op.getLoc(), rewriter, /*requireLeafArrays=*/false,
3605 /*requireShapeCarrier=*/true
3606 )
3607 .shapeCarrier;
3608 } else if (auto it = resolver.deferredPodArrays.find(readOp.getResult());
3609 it != resolver.deferredPodArrays.end() && it->second.shapeCarrier) {
3610 shapeSource = castValueToTypeIfNeeded(
3611 rewriter, op.getLoc(), it->second.shapeCarrier, getPodArrayShapeCarrierType(arrTy)
3612 );
3613 }
3614 }
3615
3616 if (!shapeSource) {
3617 shapeSource = materializeArrayLengthCarrier(op.getArrRef(), arrTy, op.getLoc(), rewriter);
3618 }
3619
3620 preserveDiscardableAttrsExcept(
3621 op, rewriter.replaceOpWithNewOp<ArrayLengthOp>(op, shapeSource, op.getDim()),
3622 DEFERRED_POD_ARRAY_LENGTH_ATTR
3623 );
3624 return success();
3625 }
3626};
3627
3629class SplitPodArrayExtractArrayOp : public OpConversionPattern<ExtractArrayOp> {
3630public:
3631 using OpConversionPattern<ExtractArrayOp>::OpConversionPattern;
3632
3633 static bool legal(ExtractArrayOp op) { return !splittablePodArray(op.getResult().getType()); }
3634
3635 LogicalResult matchAndRewrite(
3636 ExtractArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3637 ) const override {
3638 if (legal(op)) {
3639 return failure();
3640 }
3641
3642 SmallVector<Type> splitResultTypes;
3643 splitPodArrayTypeTo(op.getResult().getType(), splitResultTypes);
3644 if (splitResultTypes.empty()) {
3645 ArrayType resultTy = llvm::cast<ArrayType>(op.getResult().getType());
3646 preserveDiscardableAttrs(
3647 op, rewriter.replaceOpWithNewOp<ExtractArrayOp>(
3648 op, getPodArrayShapeCarrierType(resultTy),
3649 getSingleConvertedValue(adaptor.getArrRef()),
3650 flattenConvertedValues(adaptor.getIndices())
3651 )
3652 );
3653 return success();
3654 }
3655
3656 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
3657 SmallVector<Value> replacements;
3658 ArrayType resultTy = llvm::cast<ArrayType>(op.getResult().getType());
3659 replacements.reserve(splitResultTypes.size() + (needsPodArrayShapeCarrier(resultTy) ? 1 : 0));
3660 auto splitArrRefs = adaptor.getArrRef().take_front(splitResultTypes.size());
3661 for (auto [splitArrRange, splitResultType] : llvm::zip_equal(splitArrRefs, splitResultTypes)) {
3662 replacements.push_back(preserveDiscardableAttrs(
3663 op, rewriter.create<ExtractArrayOp>(
3664 op.getLoc(), llvm::cast<ArrayType>(splitResultType),
3665 getSingleConvertedValue(splitArrRange), indices
3666 )
3667 ));
3668 }
3669 if (needsPodArrayShapeCarrier(resultTy)) {
3670 Value shapeCarrier = materializeExtractedPodArrayShapeCarrier(
3671 op, resultTy, op.getArrRef(), adaptor.getArrRef(), indices, rewriter
3672 );
3673 preserveDiscardableAttrs(op, shapeCarrier.getDefiningOp());
3674 replacements.push_back(shapeCarrier);
3675 }
3676
3677 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
3678 return success();
3679 }
3680};
3681
3683class SplitPodArrayInsertArrayOp : public OpConversionPattern<InsertArrayOp> {
3684public:
3685 using OpConversionPattern<InsertArrayOp>::OpConversionPattern;
3686
3687 static bool legal(InsertArrayOp op) { return !splittablePodArray(op.getRvalue().getType()); }
3688
3689 LogicalResult matchAndRewrite(
3690 InsertArrayOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3691 ) const override {
3692 if (legal(op)) {
3693 return failure();
3694 }
3695
3696 if (hasZeroLeafPodArraySplit(llvm::cast<ArrayType>(op.getRvalue().getType()))) {
3697 preserveDiscardableAttrs(
3698 op, rewriter.create<InsertArrayOp>(
3699 op.getLoc(), getSingleConvertedValue(adaptor.getArrRef()),
3700 flattenConvertedValues(adaptor.getIndices()),
3701 getSingleConvertedValue(adaptor.getRvalue())
3702 )
3703 );
3704 rewriter.eraseOp(op);
3705 return success();
3706 }
3707
3708 ArrayType destArrTy = llvm::cast<ArrayType>(op.getArrRef().getType());
3709 ArrayType rvalueTy = llvm::cast<ArrayType>(op.getRvalue().getType());
3710 SmallVector<Value> indices = flattenConvertedValues(adaptor.getIndices());
3711 size_t leafCount = getSplitPodArrayLeafCount(rvalueTy);
3712 auto splitArrRefs = adaptor.getArrRef().take_front(leafCount);
3713 auto splitRvalues = adaptor.getRvalue().take_front(leafCount);
3714 for (auto [splitArrRange, splitRvalueRange] : llvm::zip_equal(splitArrRefs, splitRvalues)) {
3715 preserveDiscardableAttrs(
3716 op, rewriter.create<InsertArrayOp>(
3717 op.getLoc(), getSingleConvertedValue(splitArrRange), indices,
3718 getSingleConvertedValue(splitRvalueRange)
3719 )
3720 );
3721 }
3722 if (needsPodArrayShapeCarrier(destArrTy)) {
3723 Value destCarrier = getConvertedPodArrayShapeCarrierIfPresent(destArrTy, adaptor.getArrRef());
3724 if (!destCarrier) {
3725 return rewriter.notifyMatchFailure(
3726 op, "expected converted destination shape carrier for array-of-pod insert"
3727 );
3728 }
3729
3730 Value rvalueCarrier =
3731 getConvertedPodArrayShapeCarrierIfPresent(rvalueTy, adaptor.getRvalue());
3732 if (!rvalueCarrier) {
3733 rvalueCarrier =
3734 materializeArrayLengthCarrier(op.getRvalue(), rvalueTy, op.getLoc(), rewriter);
3735 }
3736
3737 preserveDiscardableAttrs(
3738 op, rewriter.create<InsertArrayOp>(op.getLoc(), destCarrier, indices, rvalueCarrier)
3739 );
3740 }
3741
3742 rewriter.eraseOp(op);
3743 return success();
3744 }
3745};
3746
3748class SplitPodArrayInMemberWriteOp : public OpConversionPattern<MemberWriteOp> {
3749 SymbolTableCollection &tables;
3750 const MemberReplacementMap &repMapRef;
3751
3752public:
3753 SplitPodArrayInMemberWriteOp(
3754 const TypeConverter &converter, MLIRContext *ctx, SymbolTableCollection &symTables,
3755 const MemberReplacementMap &memberRepMap
3756 )
3757 : OpConversionPattern<MemberWriteOp>(converter, ctx), tables(symTables),
3758 repMapRef(memberRepMap) {}
3759
3760 static bool legal(MemberWriteOp op) { return !splittablePodArray(op.getVal().getType()); }
3761
3762 LogicalResult matchAndRewrite(
3763 MemberWriteOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3764 ) const override {
3765 if (legal(op)) {
3766 return failure();
3767 }
3768 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
3769 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
3770 assert(succeeded(tgtStructDef));
3771
3772 const LocalMemberReplacementMap &idToMember =
3773 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
3774 ArrayType arrTy = llvm::cast<ArrayType>(op.getVal().getType());
3775 SmallVector<RecordChain> splitIds;
3776 SmallVector<Type> splitTypes;
3777 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
3778 if (splitTypes.empty()) {
3779 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3780 preserveDiscardableAttrs(
3781 op,
3782 rewriter.create<MemberWriteOp>(
3783 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3784 FlatSymbolRefAttr::get(carrierMember.first), getSingleConvertedValue(adaptor.getVal())
3785 )
3786 );
3787 rewriter.eraseOp(op);
3788 return success();
3789 }
3790
3791 auto splitVals = adaptor.getVal().take_front(splitIds.size());
3792 for (auto [id, splitValRange] : llvm::zip_equal(splitIds, splitVals)) {
3793 const MemberInfo &newMember = idToMember.at(id);
3794 preserveDiscardableAttrs(
3795 op, rewriter.create<MemberWriteOp>(
3796 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3797 FlatSymbolRefAttr::get(newMember.first), getSingleConvertedValue(splitValRange)
3798 )
3799 );
3800 }
3801 if (needsPodArrayShapeCarrier(arrTy)) {
3802 Value carrier = getConvertedPodArrayShapeCarrierIfPresent(arrTy, adaptor.getVal());
3803 if (!carrier) {
3804 carrier = materializeArrayLengthCarrier(op.getVal(), arrTy, op.getLoc(), rewriter);
3805 }
3806 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3807 preserveDiscardableAttrs(
3808 op, rewriter.create<MemberWriteOp>(
3809 op.getLoc(), getSingleConvertedValue(adaptor.getComponent()),
3810 FlatSymbolRefAttr::get(carrierMember.first),
3811 castValueToTypeIfNeeded(rewriter, op.getLoc(), carrier, carrierMember.second)
3812 )
3813 );
3814 }
3815 rewriter.eraseOp(op);
3816 return success();
3817 }
3818};
3819
3821class SplitPodArrayInMemberReadOp : public OpConversionPattern<MemberReadOp> {
3822 SymbolTableCollection &tables;
3823 const MemberReplacementMap &repMapRef;
3824
3825public:
3826 SplitPodArrayInMemberReadOp(
3827 const TypeConverter &converter, MLIRContext *ctx, SymbolTableCollection &symTables,
3828 const MemberReplacementMap &memberRepMap
3829 )
3830 : OpConversionPattern<MemberReadOp>(converter, ctx), tables(symTables),
3831 repMapRef(memberRepMap) {}
3832
3833 static bool legal(MemberReadOp op) { return !splittablePodArray(op.getResult().getType()); }
3834
3835 LogicalResult matchAndRewrite(
3836 MemberReadOp op, OneToNOpAdaptor adaptor, ConversionPatternRewriter &rewriter
3837 ) const override {
3838 if (legal(op)) {
3839 return failure();
3840 }
3841 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
3842 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
3843 assert(succeeded(tgtStructDef));
3844
3845 const LocalMemberReplacementMap &idToMember =
3846 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
3847 ArrayType arrTy = llvm::cast<ArrayType>(op.getType());
3848 SmallVector<RecordChain> splitIds;
3849 SmallVector<Type> splitTypes;
3850 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
3851 SmallVector<Value> mapOperands;
3852 std::optional<int32_t> numDimsPerMap;
3853 auto mapOperandsOld = adaptor.getMapOperands();
3854 if (!mapOperandsOld.empty()) {
3855 assert(
3856 mapOperandsOld.size() == 1 &&
3857 "member.readm should have at most one affine-map operand group"
3858 );
3859 mapOperands = flattenConvertedValues(mapOperandsOld.front());
3860
3861 ArrayRef<int32_t> numDimsPerMapOld = op.getNumDimsPerMap();
3862 if (!numDimsPerMapOld.empty()) {
3863 assert(
3864 numDimsPerMapOld.size() == 1 &&
3865 "member.readm should have one numDims entry per affine-map group"
3866 );
3867 numDimsPerMap = numDimsPerMapOld.front();
3868 }
3869 }
3870 if (splitTypes.empty()) {
3871 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3872 Value carrierRead = preserveDiscardableAttrs(
3873 op,
3874 rewriter.create<MemberReadOp>(
3875 op.getLoc(), carrierMember.second, getSingleConvertedValue(adaptor.getComponent()),
3876 carrierMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3877 )
3878 );
3879 rewriter.replaceOpWithMultiple(op, {ValueRange {carrierRead}});
3880 return success();
3881 }
3882 SmallVector<Value> replacements;
3883 replacements.reserve(splitIds.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0));
3884 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
3885 const MemberInfo &newMember = idToMember.at(id);
3886 replacements.push_back(preserveDiscardableAttrs(
3887 op, rewriter.create<MemberReadOp>(
3888 op.getLoc(), splitType, getSingleConvertedValue(adaptor.getComponent()),
3889 newMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3890 )
3891 ));
3892 }
3893 if (needsPodArrayShapeCarrier(arrTy)) {
3894 const MemberInfo &carrierMember = idToMember.at(RecordChain());
3895 replacements.push_back(preserveDiscardableAttrs(
3896 op,
3897 rewriter.create<MemberReadOp>(
3898 op.getLoc(), carrierMember.second, getSingleConvertedValue(adaptor.getComponent()),
3899 carrierMember.first, op.getTableOffset().value_or(nullptr), mapOperands, numDimsPerMap
3900 )
3901 ));
3902 }
3903 rewriter.replaceOpWithMultiple(op, {ValueRange(replacements)});
3904 return success();
3905 }
3906};
3907
3911static LogicalResult rejectUnsupportedPodArrayUnifiableCasts(ModuleOp modOp) {
3912 auto result = modOp.walk([](UnifiableCastOp op) -> WalkResult {
3913 ArrayType inputArrTy = splittablePodArray(op.getInput().getType());
3914 ArrayType resultArrTy = splittablePodArray(op.getType());
3915 if (!inputArrTy && !resultArrTy) {
3916 return WalkResult::advance();
3917 }
3918 if (!inputArrTy && resultArrTy) {
3919 return op.emitOpError(
3920 "cannot lower a generic input to a split array-of-POD result without rewriting "
3921 "the surrounding function/template signature"
3922 );
3923 }
3924 if (inputArrTy && !resultArrTy) {
3925 return op.emitOpError(
3926 "cannot lower a split array-of-POD input to a non-array result without materializing "
3927 "the aggregate or rewriting the surrounding function/template signature"
3928 );
3929 }
3930 return WalkResult::advance();
3931 });
3932 return failure(result.wasInterrupted());
3933}
3934
3936static LogicalResult
3937step2(ModuleOp modOp, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap) {
3938 if (failed(rejectUnsupportedPodArrayUnifiableCasts(modOp))) {
3939 return failure();
3940 }
3941
3942 MLIRContext *ctx = modOp.getContext();
3943 PodArrayTypeConverter typeConverter;
3944 CompatiblePodLeafMaterializationMap materializedLeaves;
3945
3946 RewritePatternSet patterns(ctx);
3947 patterns.add<
3948 SplitPodArrayNonDetOp, SplitPodArrayCreateArrayOp, SplitPodArrayReadArrayOp,
3949 SplitPodArrayWriteArrayOp, SplitPodArrayExtractArrayOp, SplitPodArrayInsertArrayOp,
3950 SplitPodArrayInFuncDefOp, SplitPodArrayInUnifiableCastOp, SplitPodArrayInReturnOp,
3951 SplitPodArrayInCallOp, RejectRaggedNestedLeafArrayLengthOp, SplitPodArrayLengthOp>(
3952 typeConverter, ctx
3953 );
3954 patterns.add<SplitPodArrayInEmitEqualityOp, SplitPodArrayInEmitContainmentOp>(
3955 typeConverter, ctx, materializedLeaves
3956 );
3957 patterns.add<SplitPodArrayInMemberWriteOp, SplitPodArrayInMemberReadOp>(
3958 typeConverter, ctx, symTables, memberRepMap
3959 );
3960
3961 ConversionTarget target(*ctx);
3962 baseTargetSetup(target);
3963 target.addLegalOp<UnrealizedConversionCastOp>();
3964 target.addDynamicallyLegalOp<NonDetOp>(SplitPodArrayNonDetOp::legal);
3965 target.addDynamicallyLegalOp<CreateArrayOp>(SplitPodArrayCreateArrayOp::legal);
3966 target.addDynamicallyLegalOp<ReadArrayOp>(SplitPodArrayReadArrayOp::legal);
3967 target.addDynamicallyLegalOp<WriteArrayOp>(SplitPodArrayWriteArrayOp::legal);
3968 target.addDynamicallyLegalOp<ExtractArrayOp>(SplitPodArrayExtractArrayOp::legal);
3969 target.addDynamicallyLegalOp<InsertArrayOp>(SplitPodArrayInsertArrayOp::legal);
3970 target.addDynamicallyLegalOp<FuncDefOp>(SplitPodArrayInFuncDefOp::legal);
3971 target.addDynamicallyLegalOp<UnifiableCastOp>(SplitPodArrayInUnifiableCastOp::legal);
3972 target.addDynamicallyLegalOp<ReturnOp>(SplitPodArrayInReturnOp::legal);
3973 target.addDynamicallyLegalOp<CallOp>(SplitPodArrayInCallOp::legal);
3974 target.addDynamicallyLegalOp<constrain::EmitEqualityOp>(SplitPodArrayInEmitEqualityOp::legal);
3975 target.addDynamicallyLegalOp<constrain::EmitContainmentOp>(
3976 SplitPodArrayInEmitContainmentOp::legal
3977 );
3978 target.addDynamicallyLegalOp<ArrayLengthOp>(SplitPodArrayLengthOp::legal);
3979 target.addDynamicallyLegalOp<MemberWriteOp>(SplitPodArrayInMemberWriteOp::legal);
3980 target.addDynamicallyLegalOp<MemberReadOp>(SplitPodArrayInMemberReadOp::legal);
3981
3982 mlir::scf::populateSCFStructuralTypeConversionsAndLegality(typeConverter, patterns, target);
3983
3984 LLVM_DEBUG(llvm::dbgs() << "Begin step 2: split arrays with POD element type\n";);
3985 return applyPartialConversion(modOp, target, std::move(patterns));
3986}
3987
3989class SplitInitFromNewPodOp : public OpConversionPattern<NewPodOp> {
3990public:
3991 using OpConversionPattern<NewPodOp>::OpConversionPattern;
3992
3993 static bool legal(NewPodOp op) { return op.getInitialValues().empty(); }
3994
3995 LogicalResult matchAndRewrite(
3996 NewPodOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
3997 ) const override {
3998 if (legal(op)) {
3999 return failure();
4000 }
4001 // Generate an individual write for each initialization
4002 rewriter.setInsertionPointAfter(op);
4003 Location loc = op.getLoc();
4004 for (auto [name, init] :
4005 llvm::zip_equal(adaptor.getInitializedRecords(), adaptor.getInitialValues())) {
4006 // Create the write
4007 rewriter.create<WritePodOp>(loc, op.getResult(), llvm::cast<StringAttr>(name), init);
4008 }
4009 // Remove initializations from `op`
4010 rewriter.modifyOpInPlace(op, [&op]() {
4011 op.getInitialValuesMutable().clear();
4012 op.setInitializedRecordsAttr(ArrayAttr::get(op.getContext(), {})); // DefaultValuedAttr:{}
4013 });
4014 return success();
4015 }
4016};
4017
4024class SplitPodElementCreateArrayOp : public OpConversionPattern<CreateArrayOp> {
4025 const Step3Resolver &resolver;
4026
4027public:
4028 SplitPodElementCreateArrayOp(MLIRContext *ctx, const Step3Resolver &step3Resolver)
4029 : OpConversionPattern<CreateArrayOp>(ctx), resolver(step3Resolver) {}
4030
4031 static bool legal(CreateArrayOp op) {
4032 return !llvm::any_of(op.getElements().getTypes(), [](Type type) {
4033 return splittablePod(type) || llvm::isa<ArrayType>(type);
4034 });
4035 }
4036
4037 LogicalResult matchAndRewrite(
4038 CreateArrayOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4039 ) const override {
4040 if (legal(op)) {
4041 return failure();
4042 }
4043 SmallVector<Value> leafElements;
4044 leafElements.reserve(adaptor.getElements().size());
4045
4046 Type leafType;
4047 for (Value element : adaptor.getElements()) {
4048 SmallVector<Value> flattenedValues;
4049 if (splittablePod(element.getType())) {
4050 processInputOperand(
4051 op.getLoc(), element, flattenedValues, rewriter, op.getOperation(),
4052 &resolver.virtualPods
4053 );
4054 } else {
4055 flattenedValues.push_back(element);
4056 }
4057
4058 assert(
4059 flattenedValues.size() == 1 &&
4060 "array.new elements should already have been split to a single flattened leaf"
4061 );
4062 if (!leafType) {
4063 leafType = flattenedValues.front().getType();
4064 } else {
4065 assert(
4066 leafType == flattenedValues.front().getType() && "array.new elements must stay uniform"
4067 );
4068 }
4069 leafElements.push_back(flattenedValues.front());
4070 }
4071
4072 size_t leafRank = 0;
4073 if (auto leafArrTy = llvm::dyn_cast_if_present<ArrayType>(leafType)) {
4074 leafRank = leafArrTy.getDimensionSizes().size();
4075 }
4076 ArrayType arrTy = op.getType();
4077 assert(
4078 arrTy.getDimensionSizes().size() >= leafRank && "flattened leaf rank exceeds array rank"
4079 );
4080 size_t outerRank = arrTy.getDimensionSizes().size() - leafRank;
4081 assert(outerRank > 0 && "array.new elements must populate at least one outer array dimension");
4082
4083 ArrayType outerIndexTy =
4084 ArrayType::get(arrTy.getElementType(), arrTy.getDimensionSizes().take_front(outerRank));
4085 auto elementIndices = outerIndexTy.getSubelementIndices();
4086 assert(
4087 elementIndices && "array.new with explicit POD elements requires static outer dimensions"
4088 );
4089 assert(
4090 elementIndices->size() == leafElements.size() &&
4091 "array.new element count must match the outer array cardinality"
4092 );
4093
4094 Value rebuiltArray = createWritableArrayValue(rewriter, op.getLoc(), arrTy);
4095 preserveDiscardableAttrs(op, rebuiltArray.getDefiningOp());
4096 for (auto [index, leafValue] : llvm::zip_equal(*elementIndices, leafElements)) {
4097 genArrayWrite(rewriter, op.getLoc(), rebuiltArray, index, leafValue);
4098 }
4099 rewriter.replaceOp(op, rebuiltArray);
4100 return success();
4101 }
4102};
4103
4111class SplitPodInFuncDefOp : public OpConversionPattern<FuncDefOp> {
4112 Step3Resolver &resolver;
4113
4114public:
4115 SplitPodInFuncDefOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4116 : OpConversionPattern<FuncDefOp>(ctx), resolver(step3Resolver) {}
4117
4118 inline static bool legal(FuncDefOp op) {
4119 return !containsSplittablePodType(op.getArgumentTypes()) &&
4120 !containsSplittablePodType(op.getResultTypes());
4121 }
4122
4123 LogicalResult
4124 matchAndRewrite(FuncDefOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
4125 if (legal(op)) {
4126 return failure();
4127 }
4128 // Update in/out types of the function to replace pods with scalars
4129 class Impl : public FunctionTypeConverter {
4130 SmallVector<size_t> originalInputIdxToSize, originalResultIdxToSize;
4131 SplitFunctionNameInfo inputNameInfo;
4132 SplitFunctionNameInfo resultNameInfo;
4133 Step3Resolver &resolver;
4134
4135 protected:
4136 SmallVector<Type> convertInputs(ArrayRef<Type> origTypes) override {
4137 return splitPodType(origTypes, &originalInputIdxToSize);
4138 }
4139 SmallVector<Type> convertResults(ArrayRef<Type> origTypes) override {
4140 return splitPodType(origTypes, &originalResultIdxToSize);
4141 }
4142 ArrayAttr convertInputAttrs(ArrayAttr origAttrs, SmallVector<Type> newTypes) override {
4144 origAttrs, originalInputIdxToSize, newTypes, ARG_NAME_ATTR_NAME,
4145 inputNameInfo.originalNames, inputNameInfo.existingNames,
4146 inputNameInfo.splitNameSuffixes
4147 );
4148 }
4149 ArrayAttr convertResultAttrs(ArrayAttr origAttrs, SmallVector<Type> newTypes) override {
4151 origAttrs, originalResultIdxToSize, newTypes, RES_NAME_ATTR_NAME,
4152 resultNameInfo.originalNames, resultNameInfo.existingNames,
4153 resultNameInfo.splitNameSuffixes
4154 );
4155 }
4156
4161 void processBlockArgs(Block &entryBlock, RewriterBase &rewriter) override {
4162 OpBuilder::InsertionGuard guard(rewriter);
4163 rewriter.setInsertionPointToStart(&entryBlock);
4164
4165 for (unsigned i = 0; i < entryBlock.getNumArguments();) {
4166 Value oldV = entryBlock.getArgument(i);
4167 if (PodType pt = splittablePod(oldV.getType())) {
4168 Location loc = oldV.getLoc();
4169 VirtualPodLeafMap leafValues;
4170 SmallVector<StringAttr> recordChain;
4171 unsigned nextArgIdx = i + 1;
4172 forEachPodLeaf(pt, recordChain, [&](const RecordChain &id, Type leafType) {
4173 BlockArgument newArg = entryBlock.insertArgument(nextArgIdx, leafType, loc);
4174 leafValues[id] = newArg;
4175 ++nextArgIdx;
4176 });
4177
4178 Value virtualPod = createVirtualPodPlaceholder(rewriter, loc, pt, leafValues);
4179 rewriter.replaceAllUsesWith(oldV, virtualPod);
4180 entryBlock.eraseArgument(i);
4181
4182 i += leafValues.size();
4183 resolver.virtualPods[virtualPod] = std::move(leafValues);
4184 } else {
4185 ++i;
4186 }
4187 }
4188 }
4189
4190 public:
4191 Impl(FuncDefOp op, Step3Resolver &step3Resolver) : resolver(step3Resolver) {
4192 inputNameInfo = collectSplitFunctionNameInfo(op.getArgumentTypes(), [&op](unsigned i) {
4193 return op.getArgNameAttr(i);
4194 }, getSplitRecordNameSuffixes);
4195 resultNameInfo = collectSplitFunctionNameInfo(
4196 op.getResultTypes(), [resultAttrs = op.getAllResultAttrs()](unsigned i) {
4197 return getAttrAtIndexWithName(resultAttrs, i, RES_NAME_ATTR_NAME);
4198 }, getSplitRecordNameSuffixes
4199 );
4200 }
4201 };
4202 Impl(op, resolver).convert(op, rewriter);
4203 return success();
4204 }
4205};
4206
4212class SplitPodInReturnOp : public OpConversionPattern<ReturnOp> {
4213 const Step3Resolver &resolver;
4214
4215public:
4216 SplitPodInReturnOp(MLIRContext *ctx, const Step3Resolver &step3Resolver)
4217 : OpConversionPattern<ReturnOp>(ctx), resolver(step3Resolver) {}
4218
4219 inline static bool legal(ReturnOp op) {
4220 return !containsSplittablePodType(op.getOperands().getTypes());
4221 }
4222
4223 LogicalResult matchAndRewrite(
4224 ReturnOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4225 ) const override {
4226 if (legal(op)) {
4227 return failure();
4228 }
4229 processInputOperands(
4230 adaptor.getOperands(), op.getOperandsMutable(), op, rewriter, &resolver.virtualPods
4231 );
4232 return success();
4233 }
4234};
4235
4237static CallOp newCallOpWithSplitResults(
4238 CallOp oldCall, CallOp::Adaptor adaptor, ConversionPatternRewriter &rewriter,
4239 Step3Resolver &resolver
4240) {
4241 OpBuilder::InsertionGuard guard(rewriter);
4242 rewriter.setInsertionPointAfter(oldCall);
4243
4244 Operation::result_range oldResults = oldCall.getResults();
4246 oldCall.getLoc(), splitPodType(oldResults.getTypes()), oldCall, adaptor.getMapOperands(),
4247 adaptor.getArgOperands(), rewriter
4248 );
4249
4250 auto newResults = newCall.getResults().begin();
4251 for (Value oldVal : oldResults) {
4252 if (PodType pt = splittablePod(oldVal.getType())) {
4253 Location loc = oldVal.getLoc();
4254 VirtualPodLeafMap leafValues;
4255 SmallVector<StringAttr> recordChain;
4256 forEachPodLeaf(pt, recordChain, [&leafValues, &newResults](const RecordChain &id, Type) {
4257 leafValues[id] = *newResults;
4258 ++newResults;
4259 });
4260 Value virtualPod = createVirtualPodPlaceholder(rewriter, loc, pt, leafValues);
4261 resolver.virtualPods[virtualPod] = std::move(leafValues);
4262 rewriter.replaceAllUsesWith(oldVal, virtualPod);
4263 } else {
4264 rewriter.replaceAllUsesWith(oldVal, *newResults);
4265 newResults++;
4266 }
4267 }
4268 // erase the original CallOp
4269 rewriter.eraseOp(oldCall);
4270
4271 return newCall;
4272}
4273
4280class SplitPodInCallOp : public OpConversionPattern<CallOp> {
4281 Step3Resolver &resolver;
4282
4283public:
4284 SplitPodInCallOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4285 : OpConversionPattern<CallOp>(ctx), resolver(step3Resolver) {}
4286
4287 inline static bool legal(CallOp op) {
4288 return !containsSplittablePodType(op.getArgOperands().getTypes()) &&
4289 !containsSplittablePodType(op.getResultTypes());
4290 }
4291
4292 LogicalResult matchAndRewrite(
4293 CallOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4294 ) const override {
4295 if (legal(op)) {
4296 return failure();
4297 }
4298 // Create new CallOp with split results first so, then process its inputs to split types
4299 CallOp newCall = newCallOpWithSplitResults(op, adaptor, rewriter, resolver);
4300 processInputOperands(
4301 newCall.getArgOperands(), newCall.getArgOperandsMutable(), newCall, rewriter,
4302 &resolver.virtualPods
4303 );
4304 return success();
4305 }
4306};
4307
4309class SplitPodInMemberWriteOp : public OpConversionPattern<MemberWriteOp> {
4310 SymbolTableCollection &tables;
4311 const MemberReplacementMap &repMapRef;
4312 const Step3Resolver &resolver;
4313
4314public:
4315 SplitPodInMemberWriteOp(
4316 MLIRContext *ctx, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap,
4317 const Step3Resolver &step3Resolver
4318 )
4319 : OpConversionPattern<MemberWriteOp>(ctx), tables(symTables), repMapRef(memberRepMap),
4320 resolver(step3Resolver) {}
4321
4322 static bool legal(MemberWriteOp op) { return !containsSplittablePodType(op.getVal().getType()); }
4323
4324 LogicalResult matchAndRewrite(
4325 MemberWriteOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4326 ) const override {
4327 if (legal(op)) {
4328 return failure();
4329 }
4330 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
4331 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
4332 assert(succeeded(tgtStructDef));
4333
4334 const LocalMemberReplacementMap &idToMember =
4335 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
4336 const VirtualPodLeafMap *virtualLeafValues =
4337 !hasEarlierWriteToPod(op.getOperation(), op.getVal())
4338 ? lookupVirtualPodLeafMap(op.getVal(), resolver.virtualPods)
4339 : nullptr;
4340
4341 for (const auto &[id, newMember] : idToMember) {
4342 Value scalarValue = virtualLeafValues
4343 ? virtualLeafValues->at(id)
4344 : genReadAlongPath(rewriter, op.getLoc(), op.getVal(), id);
4345 preserveDiscardableAttrs(
4346 op, rewriter.create<MemberWriteOp>(
4347 op.getLoc(), adaptor.getComponent(), FlatSymbolRefAttr::get(newMember.first),
4348 scalarValue
4349 )
4350 );
4351 }
4352 rewriter.eraseOp(op);
4353 return success();
4354 }
4355};
4356
4358class SplitPodInMemberReadOp : public OpConversionPattern<MemberReadOp> {
4359 SymbolTableCollection &tables;
4360 const MemberReplacementMap &repMapRef;
4361 Step3Resolver &resolver;
4362
4363public:
4364 SplitPodInMemberReadOp(
4365 MLIRContext *ctx, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap,
4366 Step3Resolver &step3Resolver
4367 )
4368 : OpConversionPattern<MemberReadOp>(ctx), tables(symTables), repMapRef(memberRepMap),
4369 resolver(step3Resolver) {}
4370
4371 static bool legal(MemberReadOp op) {
4372 return !containsSplittablePodType(op.getResult().getType());
4373 }
4374
4375 LogicalResult matchAndRewrite(
4376 MemberReadOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4377 ) const override {
4378 if (legal(op)) {
4379 return failure();
4380 }
4381 StructType tgtStructTy = llvm::cast<MemberRefOpInterface>(op.getOperation()).getStructType();
4382 auto tgtStructDef = tgtStructTy.getDefinition(tables, op);
4383 assert(succeeded(tgtStructDef));
4384
4385 const LocalMemberReplacementMap &idToMember =
4386 repMapRef.at(tgtStructDef->get()).at(op.getMemberNameAttr().getAttr());
4387
4388 VirtualPodLeafMap leafValues;
4389 for (const auto &[id, newMember] : idToMember) {
4390 leafValues[id] = preserveDiscardableAttrs(
4391 op, rewriter.create<MemberReadOp>(
4392 op.getLoc(), newMember.second, adaptor.getComponent(), newMember.first
4393 )
4394 );
4395 }
4396
4397 PodType podTy = llvm::cast<PodType>(op.getType());
4398 Value virtualPod = createVirtualPodPlaceholder(rewriter, op.getLoc(), podTy, leafValues);
4399 resolver.virtualPods[virtualPod] = std::move(leafValues);
4400 rewriter.replaceOp(op, virtualPod);
4401 return success();
4402 }
4403};
4404
4409static bool tryCollectMaterializedSplitPodArrayLeafValues(
4410 Value arrayValue, ArrayType arrTy, ArrayRef<Type> splitTypes, SmallVectorImpl<Value> &leafArrays
4411) {
4412 auto cast = arrayValue.getDefiningOp<UnrealizedConversionCastOp>();
4413 size_t expectedOperands = splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0);
4414 if (!cast || cast->getNumResults() != 1 || cast.getResult(0).getType() != arrTy ||
4415 cast->getNumOperands() != expectedOperands) {
4416 return false;
4417 }
4418 if (needsPodArrayShapeCarrier(arrTy) &&
4419 cast.getOperand(splitTypes.size()).getType() != getPodArrayShapeCarrierType(arrTy)) {
4420 return false;
4421 }
4422
4424 cast.getOperands().take_front(splitTypes.size()), splitTypes, leafArrays
4425 );
4426}
4427
4433static bool tryCollectReadPodSplitPodArrayLeafValues(
4434 ReadPodOp readOp, ArrayType arrTy, ArrayRef<RecordChain> splitIds, ArrayRef<Type> splitTypes,
4435 const VirtualPodValueMap &virtualPods, SmallVectorImpl<Value> &leafArrays
4436) {
4437 auto tryCollectFromVirtualRead = [&](ReadPodOp sourceRead) {
4438 if (hasEarlierWrite(sourceRead)) {
4439 return false;
4440 }
4441
4442 const VirtualPodLeafMap *podLeafValues =
4443 lookupVirtualPodLeafMap(sourceRead.getPodRef(), virtualPods);
4444 if (!podLeafValues) {
4445 return false;
4446 }
4447
4448 SmallVector<Value> stagedLeafArrays;
4449 stagedLeafArrays.reserve(splitIds.size());
4450 for (const RecordChain &id : splitIds) {
4451 auto it = podLeafValues->find(id.withPrefix({sourceRead.getRecordNameAttr()}));
4452 if (it == podLeafValues->end() ||
4453 !typesUnify(it->second.getType(), getFlattenedTypeAlongPath(arrTy, id))) {
4454 return false;
4455 }
4456 stagedLeafArrays.push_back(it->second);
4457 }
4458 llvm::append_range(leafArrays, stagedLeafArrays);
4459 return true;
4460 };
4461
4462 if (WritePodOp writeOp = findNearestForwardableWrite(readOp)) {
4463 if (tryCollectMaterializedSplitPodArrayLeafValues(
4464 writeOp.getValue(), arrTy, splitTypes, leafArrays
4465 )) {
4466 return true;
4467 }
4468
4469 if (ReadPodOp writtenRead = peelUnifiableCasts(writeOp.getValue()).getDefiningOp<ReadPodOp>()) {
4470 if (tryCollectFromVirtualRead(writtenRead)) {
4471 return true;
4472 }
4473 }
4474 }
4475
4476 if (tryCollectFromVirtualRead(readOp)) {
4477 return true;
4478 }
4479
4480 return false;
4481}
4482
4484static bool resolveReadPodSplitPodArrayLeafValues(
4485 ReadPodOp readOp, ArrayType arrTy, ArrayRef<RecordChain> splitIds, ArrayRef<Type> splitTypes,
4486 const VirtualPodValueMap &virtualPods, DeferredPodArrayBackingMap &deferredPodArrays,
4487 Location loc, OpBuilder &bldr, SmallVectorImpl<Value> &leafArrays
4488) {
4489 if (tryCollectReadPodSplitPodArrayLeafValues(
4490 readOp, arrTy, splitIds, splitTypes, virtualPods, leafArrays
4491 )) {
4492 return true;
4493 }
4494
4495 if (!isFreshUnwrittenPodRead(readOp)) {
4496 return false;
4497 }
4498
4499 // Reuse one synthetic split-array backing per deferred field read so repeated users of the same
4500 // aggregate value continue to observe the same unwritten leaf storage and shared shape witness.
4501 DeferredPodArrayBacking &backing = materializeDeferredPodArrayBacking(
4502 readOp, arrTy, splitTypes, deferredPodArrays, loc, bldr,
4503 /*requireLeafArrays=*/true, /*requireShapeCarrier=*/true
4504 );
4505 leafArrays.assign(backing.leafArrays.begin(), backing.leafArrays.end());
4506
4507 return true;
4508}
4509
4511static Value tryResolveReadPodArrayShapeSource(
4512 ReadPodOp readOp, ArrayType arrTy, const VirtualPodValueMap &virtualPods, Location loc,
4513 RewriterBase &rewriter
4514) {
4515 auto tryGetVirtualShapeSource = [&](ReadPodOp sourceRead) -> Value {
4516 if (hasEarlierWrite(sourceRead)) {
4517 return {};
4518 }
4519
4520 const VirtualPodLeafMap *podLeafValues =
4521 lookupVirtualPodLeafMap(sourceRead.getPodRef(), virtualPods);
4522 if (!podLeafValues) {
4523 return {};
4524 }
4525
4526 SmallVector<StringAttr> carrierPath {
4527 sourceRead.getRecordNameAttr(), getPodArrayShapeCarrierMarker(rewriter.getContext())
4528 };
4529 if (auto carrierIt = podLeafValues->find(RecordChain(carrierPath, true));
4530 carrierIt != podLeafValues->end()) {
4531 return castValueToTypeIfNeeded(
4532 rewriter, loc, carrierIt->second, getPodArrayShapeCarrierType(arrTy)
4533 );
4534 }
4535
4536 size_t originalRank = arrTy.getDimensionSizes().size();
4537 SmallVector<RecordChain> splitIds;
4538 SmallVector<Type> splitTypes;
4539 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4540 for (auto [id, splitType] : llvm::zip_equal(splitIds, splitTypes)) {
4541 auto splitArrTy = llvm::dyn_cast<ArrayType>(splitType);
4542 if (!splitArrTy || splitArrTy.getDimensionSizes().size() != originalRank) {
4543 continue;
4544 }
4545
4546 auto valueIt = podLeafValues->find(id.withPrefix({sourceRead.getRecordNameAttr()}));
4547 if (valueIt != podLeafValues->end()) {
4548 return castValueToTypeIfNeeded(rewriter, loc, valueIt->second, splitArrTy);
4549 }
4550 }
4551
4552 return {};
4553 };
4554
4555 if (!hasEarlierWrite(readOp)) {
4556 if (Value shapeSource = tryGetVirtualShapeSource(readOp)) {
4557 return shapeSource;
4558 }
4559 }
4560
4561 if (WritePodOp writeOp = findNearestForwardableWrite(readOp)) {
4562 if (Value shapeSource = tryCollectDirectConvertedPodArrayShapeSource(
4563 writeOp.getValue(), arrTy, loc, rewriter
4564 )) {
4565 return shapeSource;
4566 }
4567 if (ReadPodOp writtenRead = peelUnifiableCasts(writeOp.getValue()).getDefiningOp<ReadPodOp>()) {
4568 if (Value shapeSource = tryGetVirtualShapeSource(writtenRead)) {
4569 return shapeSource;
4570 }
4571 }
4572 }
4573
4574 return {};
4575}
4576
4578static void eraseDeadDeferredFieldReadChain(ReadPodOp readOp, PatternRewriter &rewriter) {
4579 if (!readOp.getResult().use_empty()) {
4580 return;
4581 }
4582
4583 Value podRef = readOp.getPodRef();
4584 rewriter.eraseOp(readOp);
4585 if (podRef.use_empty()) {
4586 if (auto cast = podRef.getDefiningOp<UnrealizedConversionCastOp>()) {
4587 if (cast->getNumResults() == 1 && cast.getResult(0) == podRef) {
4588 rewriter.eraseOp(cast);
4589 }
4590 }
4591 }
4592}
4593
4595static bool getDeferredSplitPodArrayCastInfo(
4596 UnrealizedConversionCastOp op, ArrayType &arrTy, SmallVector<RecordChain> &splitIds,
4597 SmallVectorImpl<Type> &splitTypes
4598) {
4599 if (op->getNumOperands() != 1) {
4600 return false;
4601 }
4602
4603 arrTy = splittablePodArray(op.getOperand(0).getType());
4604 if (!arrTy) {
4605 return false;
4606 }
4607
4608 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4609 size_t expectedResults = splitTypes.size() + (needsPodArrayShapeCarrier(arrTy) ? 1 : 0);
4610 if (op->getNumResults() != expectedResults) {
4611 return false;
4612 }
4613
4614 for (auto [result, splitType] :
4615 llvm::zip_equal(op.getResults().take_front(splitTypes.size()), splitTypes)) {
4616 if (result.getType() != splitType) {
4617 return false;
4618 }
4619 }
4620 return !needsPodArrayShapeCarrier(arrTy) ||
4621 op.getResult(splitTypes.size()).getType() == getPodArrayShapeCarrierType(arrTy);
4622}
4623
4625static LogicalResult resolveDeferredSplitPodArrayCast(
4626 UnrealizedConversionCastOp op, PatternRewriter &rewriter, Step3Resolver &resolver
4627) {
4628 ArrayType arrTy;
4629 SmallVector<RecordChain> splitIds;
4630 SmallVector<Type> splitTypes;
4631 if (!getDeferredSplitPodArrayCastInfo(op, arrTy, splitIds, splitTypes)) {
4632 return failure();
4633 }
4634
4635 ReadPodOp fieldRead = peelUnifiableCasts(op.getOperand(0)).getDefiningOp<ReadPodOp>();
4636 if (!fieldRead) {
4637 return failure();
4638 }
4639
4640 SmallVector<Value> splitLeafArrays;
4641 if (!resolveReadPodSplitPodArrayLeafValues(
4642 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, resolver.deferredPodArrays,
4643 op.getLoc(), rewriter, splitLeafArrays
4644 )) {
4645 return failure();
4646 }
4647
4648 SmallVector<Value> replacements(splitLeafArrays.begin(), splitLeafArrays.end());
4649 if (needsPodArrayShapeCarrier(arrTy)) {
4650 Value carrier = tryResolveReadPodArrayShapeSource(
4651 fieldRead, arrTy, resolver.virtualPods, op.getLoc(), rewriter
4652 );
4653 if (!carrier) {
4654 if (auto it = resolver.deferredPodArrays.find(fieldRead.getResult());
4655 it != resolver.deferredPodArrays.end() && it->second.shapeCarrier) {
4656 carrier = castValueToTypeIfNeeded(
4657 rewriter, op.getLoc(), it->second.shapeCarrier, getPodArrayShapeCarrierType(arrTy)
4658 );
4659 } else {
4660 carrier =
4661 materializeArrayLengthCarrier(fieldRead.getResult(), arrTy, op.getLoc(), rewriter);
4662 }
4663 }
4664 replacements.push_back(carrier);
4665 }
4666
4667 rewriter.replaceOp(op, replacements);
4668 eraseDeadDeferredFieldReadChain(fieldRead, rewriter);
4669 return success();
4670}
4671
4677class ResolvePodReadBackedArrayReadOp : public OpConversionPattern<ReadArrayOp> {
4678 Step3Resolver &resolver;
4679
4680public:
4681 ResolvePodReadBackedArrayReadOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4682 : OpConversionPattern<ReadArrayOp>(ctx), resolver(step3Resolver) {}
4683
4684 static bool canResolve(ReadArrayOp op, const Step3Resolver &resolver) {
4685 if (!shouldDeferPodArrayReadToStep3(op)) {
4686 return false;
4687 }
4688
4689 ArrayType arrTy = op.getArrRefType();
4690 auto fieldRead = llvm::cast<ReadPodOp>(op.getArrRef().getDefiningOp());
4691 SmallVector<RecordChain> splitIds;
4692 SmallVector<Type> splitTypes;
4693 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4694
4695 SmallVector<Value> ignoredLeafArrays;
4696 return tryCollectReadPodSplitPodArrayLeafValues(
4697 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, ignoredLeafArrays
4698 ) ||
4699 isFreshUnwrittenPodRead(fieldRead);
4700 }
4701
4702 LogicalResult matchAndRewrite(
4703 ReadArrayOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4704 ) const override {
4705 auto fieldRead = op.getArrRef().getDefiningOp<ReadPodOp>();
4706 if (!fieldRead) {
4707 return failure();
4708 }
4709
4710 ArrayType arrTy = op.getArrRefType();
4711 PodType podTy = llvm::cast<PodType>(arrTy.getElementType());
4712 SmallVector<RecordChain> splitIds;
4713 SmallVector<Type> splitTypes;
4714 splitPodArrayTypeTo(arrTy, splitTypes, &splitIds);
4715
4716 SmallVector<Value> splitLeafArrays;
4717 if (!resolveReadPodSplitPodArrayLeafValues(
4718 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods,
4719 resolver.deferredPodArrays, op.getLoc(), rewriter, splitLeafArrays
4720 )) {
4721 return failure();
4722 }
4723
4724 SmallVector<Value> indices(adaptor.getIndices().begin(), adaptor.getIndices().end());
4725 VirtualPodLeafMap leafValues;
4726 for (auto [id, splitType, leafArray] : llvm::zip_equal(splitIds, splitTypes, splitLeafArrays)) {
4727 Value leafValue = ArrayAccessOpInterface::genRead(rewriter, op.getLoc(), leafArray, indices);
4728 preserveDiscardableAttrs(op, leafValue.getDefiningOp());
4729 leafValues[id] = tagRaggedNestedLeafValue(
4730 rewriter, op.getLoc(), leafValue, getRaggedNestedLeafAttrName(arrTy, splitType)
4731 );
4732 }
4733
4734 Value virtualPod = createVirtualPodPlaceholder(rewriter, op.getLoc(), podTy, leafValues);
4735 resolver.virtualPods[virtualPod] = std::move(leafValues);
4736 rewriter.replaceOp(op, virtualPod);
4737 eraseDeadDeferredFieldReadChain(fieldRead, rewriter);
4738 return success();
4739 }
4740};
4741
4749class ResolveDeferredSplitPodArrayCastOp : public OpConversionPattern<UnrealizedConversionCastOp> {
4750 Step3Resolver &resolver;
4751
4752public:
4753 ResolveDeferredSplitPodArrayCastOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4754 : OpConversionPattern<UnrealizedConversionCastOp>(ctx), resolver(step3Resolver) {}
4755
4756 static bool canResolve(UnrealizedConversionCastOp op, const Step3Resolver &resolver) {
4757 ArrayType arrTy;
4758 SmallVector<RecordChain> splitIds;
4759 SmallVector<Type> splitTypes;
4760 if (!getDeferredSplitPodArrayCastInfo(op, arrTy, splitIds, splitTypes)) {
4761 return false;
4762 }
4763
4764 ReadPodOp fieldRead = peelUnifiableCasts(op.getOperand(0)).getDefiningOp<ReadPodOp>();
4765 if (!fieldRead) {
4766 return false;
4767 }
4768
4769 SmallVector<Value> ignoredLeafArrays;
4770 return tryCollectReadPodSplitPodArrayLeafValues(
4771 fieldRead, arrTy, splitIds, splitTypes, resolver.virtualPods, ignoredLeafArrays
4772 ) ||
4773 isFreshUnwrittenPodRead(fieldRead);
4774 }
4775
4776 LogicalResult matchAndRewrite(
4777 UnrealizedConversionCastOp op, OpAdaptor, ConversionPatternRewriter &rewriter
4778 ) const override {
4779 return resolveDeferredSplitPodArrayCast(op, rewriter, resolver);
4780 }
4781};
4782
4784class ResolveDeferredSplitPodArrayCastPrepass final
4785 : public OpRewritePattern<UnrealizedConversionCastOp> {
4786 Step3Resolver &resolver;
4787
4788public:
4789 ResolveDeferredSplitPodArrayCastPrepass(MLIRContext *ctx, Step3Resolver &step3Resolver)
4790 : OpRewritePattern<UnrealizedConversionCastOp>(ctx), resolver(step3Resolver) {}
4791
4792 LogicalResult
4793 matchAndRewrite(UnrealizedConversionCastOp op, PatternRewriter &rewriter) const override {
4794 return resolveDeferredSplitPodArrayCast(op, rewriter, resolver);
4795 }
4796};
4797
4799static LogicalResult splitVirtualPodEmitEquality(
4800 constrain::EmitEqualityOp op, RewriterBase &rewriter, Step3Resolver &resolver
4801) {
4802 return splitWholePodEmitEquality(
4803 op, rewriter, resolver.materializedLeaves, op.getOperation(), &resolver.virtualPods
4804 );
4805}
4806
4813static LogicalResult splitEarlierVirtualPodEqualitiesBeforeWriteInBlock(
4814 WritePodOp writeOp, RewriterBase &rewriter, Step3Resolver &resolver
4815) {
4816 Value writtenPod = peelVirtualPodCompatibilityCasts(writeOp.getPodRef());
4817
4818 SmallVector<constrain::EmitEqualityOp> equalityOps;
4819 for (Operation &candidate : llvm::make_early_inc_range(*writeOp->getBlock())) {
4820 if (&candidate == writeOp) {
4821 break;
4822 }
4823 candidate.walk<WalkOrder::PreOrder>([&equalityOps](constrain::EmitEqualityOp op) {
4824 if (getWholePodEqualityType(op)) {
4825 equalityOps.push_back(op);
4826 }
4827 });
4828 }
4829
4830 for (constrain::EmitEqualityOp equalityOp : equalityOps) {
4831 Value lhs = peelVirtualPodCompatibilityCasts(equalityOp.getLhs());
4832 Value rhs = peelVirtualPodCompatibilityCasts(equalityOp.getRhs());
4833 if (lhs != writtenPod && rhs != writtenPod) {
4834 continue;
4835 }
4836
4837 if (!lookupVirtualPodLeafMap(equalityOp.getLhs(), resolver.virtualPods) &&
4838 !lookupVirtualPodLeafMap(equalityOp.getRhs(), resolver.virtualPods)) {
4839 continue;
4840 }
4841
4842 rewriter.setInsertionPoint(equalityOp);
4843 if (failed(splitVirtualPodEmitEquality(equalityOp, rewriter, resolver))) {
4844 return failure();
4845 }
4846 }
4847 return success();
4848}
4849
4854class ResolveVirtualPodWriteOp : public OpConversionPattern<WritePodOp> {
4855 Step3Resolver &resolver;
4856
4857public:
4858 ResolveVirtualPodWriteOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4859 : OpConversionPattern<WritePodOp>(ctx), resolver(step3Resolver) {}
4860
4861 LogicalResult matchAndRewrite(
4862 WritePodOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter
4863 ) const override {
4864 auto it = lookupVirtualPodLeafMapIt(op.getPodRef(), resolver.virtualPods);
4865 if (it == resolver.virtualPods.end() || isInsideSupportedScfRegion(op.getOperation())) {
4866 return failure();
4867 }
4868
4869 if (failed(splitEarlierVirtualPodEqualitiesBeforeWriteInBlock(op, rewriter, resolver))) {
4870 return failure();
4871 }
4872
4873 Type recordType =
4874 llvm::cast<PodType>(op.getPodRefType()).getRecordMap().lookup(op.getRecordName());
4875 assert(recordType && "record must exist in POD type");
4876 updateVirtualPodRecordLeafValues(
4877 op.getLoc(), op.getRecordNameAttr(), recordType, adaptor.getValue(), resolver.virtualPods,
4878 rewriter, it->second
4879 );
4880 rewriter.eraseOp(op);
4881 return success();
4882 }
4883};
4884
4889class ResolveVirtualPodReadOp : public OpConversionPattern<ReadPodOp> {
4890 Step3Resolver &resolver;
4891
4892public:
4893 ResolveVirtualPodReadOp(MLIRContext *ctx, Step3Resolver &step3Resolver)
4894 : OpConversionPattern<ReadPodOp>(ctx), resolver(step3Resolver) {}
4895
4896 LogicalResult
4897 matchAndRewrite(ReadPodOp op, OpAdaptor, ConversionPatternRewriter &rewriter) const override {
4898 if (hasEarlierWrite(op) || findNearestForwardableWrite(op)) {
4899 return failure();
4900 }
4901
4902 const VirtualPodLeafMap *leafValues =
4903 lookupVirtualPodLeafMap(op.getPodRef(), resolver.virtualPods);
4904 if (!leafValues) {
4905 return failure();
4906 }
4907
4908 SmallVector<StringAttr> prefix {op.getRecordNameAttr()};
4909 Type recordType =
4910 llvm::cast<PodType>(op.getPodRefType()).getRecordMap().lookup(op.getRecordName());
4911 assert(recordType && "record must exist in POD type");
4912
4913 if (PodType nestedPodTy = llvm::dyn_cast<PodType>(recordType)) {
4914 VirtualPodLeafMap nestedLeafValues;
4915 SmallVector<StringAttr> nestedRecordChain;
4916 forEachPodLeaf(nestedPodTy, nestedRecordChain, [&](const RecordChain &id, Type) {
4917 nestedLeafValues[id] = leafValues->at(id.withPrefix(prefix));
4918 });
4919 Value virtualPod =
4920 createVirtualPodPlaceholder(rewriter, op.getLoc(), nestedPodTy, nestedLeafValues);
4921 resolver.virtualPods[virtualPod] = std::move(nestedLeafValues);
4922 rewriter.replaceOp(op, virtualPod);
4923 return success();
4924 }
4925
4926 if (splittablePodArray(recordType)) {
4927 return failure();
4928 }
4929
4930 rewriter.replaceOp(
4931 op, castValueToTypeIfNeeded(
4932 rewriter, op.getLoc(), leafValues->at(RecordChain(prefix)), recordType
4933 )
4934 );
4935 return success();
4936 }
4937};
4938
4940static void rehydrateVirtualPodPlaceholders(ModuleOp modOp, VirtualPodValueMap &virtualPods) {
4941 modOp.walk([&virtualPods](UnrealizedConversionCastOp castOp) {
4942 if (castOp->getNumResults() != 1) {
4943 return;
4944 }
4945
4946 PodType podTy = llvm::dyn_cast<PodType>(castOp.getResult(0).getType());
4947 if (!podTy) {
4948 return;
4949 }
4950
4951 VirtualPodLeafMap leafValues;
4952 SmallVector<StringAttr> recordChain;
4953 auto operandIt = castOp.getOperands().begin();
4954 bool matchesVirtualPlaceholder = true;
4955 forEachPodLeaf(podTy, recordChain, [&](const RecordChain &id, Type leafType) {
4956 if (!matchesVirtualPlaceholder || operandIt == castOp.getOperands().end() ||
4957 !typesUnify((*operandIt).getType(), leafType)) {
4958 matchesVirtualPlaceholder = false;
4959 return;
4960 }
4961 leafValues[id] = *operandIt++;
4962 });
4963
4964 if (matchesVirtualPlaceholder && operandIt == castOp.getOperands().end()) {
4965 virtualPods[castOp.getResult(0)] = std::move(leafValues);
4966 }
4967 });
4968}
4969
4971class SplitVirtualPodInEmitEqualityPattern final
4972 : public OpRewritePattern<constrain::EmitEqualityOp> {
4973 Step3Resolver &resolver;
4974
4975public:
4976 SplitVirtualPodInEmitEqualityPattern(MLIRContext *ctx, Step3Resolver &step3Resolver)
4977 : OpRewritePattern<constrain::EmitEqualityOp>(ctx), resolver(step3Resolver) {}
4978
4979 LogicalResult
4980 matchAndRewrite(constrain::EmitEqualityOp op, PatternRewriter &rewriter) const override {
4981 if (!getWholePodEqualityType(op)) {
4982 return failure();
4983 }
4984
4985 return splitVirtualPodEmitEquality(op, rewriter, resolver);
4986 }
4987};
4988
4996struct PromotedFunctionArgCast {
4997 unsigned argIndex;
4998 SmallVector<Type> resultTypes;
4999 SmallVector<UnrealizedConversionCastOp> casts;
5000};
5001
5007struct PromotedFunctionSignature {
5008 FuncDefOp func;
5009 Operation *funcOp;
5010 unsigned oldInputCount;
5011 SmallVector<PromotedFunctionArgCast> argCasts;
5012};
5013
5014static const PromotedFunctionArgCast *
5015findPromotedArgCast(ArrayRef<PromotedFunctionArgCast> argCasts, unsigned argIndex) {
5016 const auto *it = llvm::find_if(argCasts, [argIndex](const PromotedFunctionArgCast &argCast) {
5017 return argCast.argIndex == argIndex;
5018 });
5019 return it == argCasts.end() ? nullptr : &*it;
5020}
5021
5028static std::optional<PromotedFunctionArgCast> getPromotableFunctionArgCast(BlockArgument arg) {
5029 PromotedFunctionArgCast promoted;
5030 promoted.argIndex = arg.getArgNumber();
5031 bool initialized = false;
5032
5033 for (OpOperand &use : arg.getUses()) {
5034 auto castOp = dyn_cast<UnrealizedConversionCastOp>(use.getOwner());
5035 if (!castOp || castOp->getNumOperands() != 1 || castOp.getOperand(0) != arg ||
5036 castOp->getNumResults() <= 1) {
5037 return std::nullopt;
5038 }
5039
5040 SmallVector<Type> castResultTypes(castOp.getResultTypes());
5041 if (!initialized) {
5042 promoted.resultTypes = std::move(castResultTypes);
5043 initialized = true;
5044 } else if (promoted.resultTypes != castResultTypes) {
5045 return std::nullopt;
5046 }
5047 promoted.casts.push_back(castOp);
5048 }
5049
5050 return initialized ? std::optional<PromotedFunctionArgCast>(std::move(promoted)) : std::nullopt;
5051}
5052
5058static ArrayAttr expandArgAttrsForPromotedFunctionArgCasts(
5059 FuncDefOp func, ArrayRef<PromotedFunctionArgCast> argCasts
5060) {
5061 ArrayAttr origAttrs = func.getArgAttrsAttr();
5062 if (!origAttrs) {
5063 return nullptr;
5064 }
5065
5066 llvm::StringSet<> usedNames;
5067 for (auto [i, attr] : llvm::enumerate(origAttrs)) {
5068 if (findPromotedArgCast(argCasts, i)) {
5069 continue;
5070 }
5071 auto dictAttr = dyn_cast<DictionaryAttr>(attr);
5072 if (!dictAttr) {
5073 continue;
5074 }
5075 if (auto nameAttr = dyn_cast_if_present<StringAttr>(dictAttr.get(ARG_NAME_ATTR_NAME))) {
5076 usedNames.insert(nameAttr.getValue());
5077 }
5078 }
5079
5080 SmallVector<Attribute> newAttrs;
5081 for (auto [i, attr] : llvm::enumerate(origAttrs)) {
5082 const PromotedFunctionArgCast *argCast = findPromotedArgCast(argCasts, i);
5083 if (!argCast) {
5084 newAttrs.push_back(attr);
5085 continue;
5086 }
5087
5088 auto dictAttr = llvm::cast<DictionaryAttr>(attr);
5089 auto nameAttr = dyn_cast_if_present<StringAttr>(dictAttr.get(ARG_NAME_ATTR_NAME));
5090 if (!nameAttr) {
5091 newAttrs.append(argCast->resultTypes.size(), attr);
5092 continue;
5093 }
5094
5095 llvm::StringRef baseName = nameAttr.getValue();
5096 for (unsigned splitIdx = 0, e = argCast->resultTypes.size(); splitIdx < e; ++splitIdx) {
5097 std::string desiredName =
5098 splitIdx == 0 ? baseName.str() : (baseName + "#" + llvm::Twine(splitIdx)).str();
5099 newAttrs.push_back(
5100 withFunctionArgNameAttr(dictAttr, reserveUniqueAttrName(usedNames, desiredName))
5101 );
5102 }
5103 }
5104 return ArrayAttr::get(func.getContext(), newAttrs);
5105}
5106
5115static std::optional<PromotedFunctionSignature> promoteFunctionArgCasts(FuncDefOp func) {
5116 if (func.isExternal()) {
5117 return std::nullopt;
5118 }
5119
5120 Block &entryBlock = func.getBody().front();
5121 SmallVector<PromotedFunctionArgCast> argCasts;
5122 for (BlockArgument arg : entryBlock.getArguments()) {
5123 if (auto promoted = getPromotableFunctionArgCast(arg)) {
5124 argCasts.push_back(std::move(*promoted));
5125 }
5126 }
5127 if (argCasts.empty()) {
5128 return std::nullopt;
5129 }
5130
5131 FunctionType oldFuncTy = func.getFunctionType();
5132 SmallVector<Type> newInputs;
5133 for (auto [i, inputType] : llvm::enumerate(oldFuncTy.getInputs())) {
5134 if (const PromotedFunctionArgCast *argCast = findPromotedArgCast(argCasts, i)) {
5135 llvm::append_range(newInputs, argCast->resultTypes);
5136 } else {
5137 newInputs.push_back(inputType);
5138 }
5139 }
5140
5141 ArrayAttr newArgAttrs = expandArgAttrsForPromotedFunctionArgCasts(func, argCasts);
5142 func.setFunctionType(FunctionType::get(func.getContext(), newInputs, oldFuncTy.getResults()));
5143 if (newArgAttrs) {
5144 func.setArgAttrsAttr(newArgAttrs);
5145 }
5146
5147 for (const PromotedFunctionArgCast &argCast : llvm::reverse(argCasts)) {
5148 BlockArgument oldArg = entryBlock.getArgument(argCast.argIndex);
5149 SmallVector<BlockArgument> newArgs;
5150 newArgs.reserve(argCast.resultTypes.size());
5151 unsigned nextArgIdx = argCast.argIndex + 1;
5152 for (Type resultType : argCast.resultTypes) {
5153 newArgs.push_back(entryBlock.insertArgument(nextArgIdx, resultType, oldArg.getLoc()));
5154 ++nextArgIdx;
5155 }
5156
5157 for (UnrealizedConversionCastOp castOp : argCast.casts) {
5158 for (auto [result, newArg] : llvm::zip_equal(castOp.getResults(), newArgs)) {
5159 result.replaceAllUsesWith(newArg);
5160 }
5161 castOp.erase();
5162 }
5163 entryBlock.eraseArgument(argCast.argIndex);
5164 }
5165
5166 return PromotedFunctionSignature {
5167 .func = func,
5168 .funcOp = func.getOperation(),
5169 .oldInputCount = oldFuncTy.getNumInputs(),
5170 .argCasts = std::move(argCasts)
5171 };
5172}
5173
5180static LogicalResult updateCallsForPromotedFunctionArgCasts(
5181 ModuleOp modOp, SymbolTableCollection &symTables,
5182 ArrayRef<PromotedFunctionSignature> promotedSignatures
5183) {
5184 if (promotedSignatures.empty()) {
5185 return success();
5186 }
5187
5188 DenseMap<Operation *, const PromotedFunctionSignature *> promotedByFunc;
5189 for (const PromotedFunctionSignature &signature : promotedSignatures) {
5190 promotedByFunc[signature.funcOp] = &signature;
5191 }
5192
5193 OpBuilder builder(modOp.getContext());
5194 for (CallOp callOp : walkCollect<CallOp>(modOp)) {
5195 FailureOr<SymbolLookupResult<FuncDefOp>> targetRes = callOp.getCalleeTarget(symTables);
5196 if (failed(targetRes)) {
5197 return failure();
5198 }
5199
5200 auto signatureIt = promotedByFunc.find(targetRes->get().getOperation());
5201 if (signatureIt == promotedByFunc.end()) {
5202 continue;
5203 }
5204
5205 const PromotedFunctionSignature &signature = *signatureIt->second;
5206 if (callOp.getArgOperands().size() != signature.oldInputCount) {
5207 return callOp.emitOpError("argument count does not match pre-promotion callee signature");
5208 }
5209
5210 SmallVector<Value> newOperands;
5211 builder.setInsertionPoint(callOp);
5212 for (auto [i, operand] : llvm::enumerate(callOp.getArgOperands())) {
5213 if (const PromotedFunctionArgCast *argCast = findPromotedArgCast(signature.argCasts, i)) {
5214 auto castOp = builder.create<UnrealizedConversionCastOp>(
5215 callOp.getLoc(), TypeRange(argCast->resultTypes), operand
5216 );
5217 llvm::append_range(newOperands, castOp.getResults());
5218 } else {
5219 newOperands.push_back(operand);
5220 }
5221 }
5222 callOp.getArgOperandsMutable().assign(ValueRange(newOperands));
5223 }
5224 return success();
5225}
5226
5232static LogicalResult
5233promoteFunctionArgCastsToSignature(ModuleOp modOp, SymbolTableCollection &symTables) {
5234 SmallVector<FuncDefOp> funcs = walkCollect<FuncDefOp>(modOp);
5235 for (unsigned promotionRounds = 1;; ++promotionRounds) {
5236 SmallVector<PromotedFunctionSignature> promotedSignatures;
5237 for (FuncDefOp func : funcs) {
5238 if (auto promoted = promoteFunctionArgCasts(func)) {
5239 promotedSignatures.push_back(std::move(*promoted));
5240 }
5241 }
5242 if (promotedSignatures.empty()) {
5243 return success();
5244 }
5245 if (failed(updateCallsForPromotedFunctionArgCasts(modOp, symTables, promotedSignatures))) {
5246 return failure();
5247 }
5248 if (promotionRounds % 64 == 0) {
5249 llvm::outs() << "function argument cast promotion has run " << promotionRounds
5250 << " rounds without reaching a fixpoint; continuing...\n";
5251 }
5252 }
5253}
5254
5255void Step3Resolver::rehydrateVirtualPodPlaceholders(ModuleOp modOp) {
5256 ::rehydrateVirtualPodPlaceholders(modOp, virtualPods);
5257}
5258
5259void Step3Resolver::addPreConversionPatterns(RewritePatternSet &patterns) {
5260 patterns.add<ResolveDeferredSplitPodArrayCastPrepass>(patterns.getContext(), *this);
5261}
5262
5263void Step3Resolver::addConversionPatterns(
5264 RewritePatternSet &patterns, SymbolTableCollection &symTables,
5265 const MemberReplacementMap &memberRepMap
5266) {
5267 patterns.add<SplitInitFromNewPodOp>(patterns.getContext());
5268 patterns.add<SplitPodElementCreateArrayOp>(patterns.getContext(), *this);
5269 patterns.add<SplitPodInFuncDefOp, SplitPodInReturnOp, SplitPodInCallOp>(
5270 patterns.getContext(), *this
5271 );
5272 patterns.add<SplitPodInMemberWriteOp, SplitPodInMemberReadOp>(
5273 patterns.getContext(), symTables, memberRepMap, *this
5274 );
5275 patterns.add<RejectRaggedNestedLeafArrayLengthOp>(patterns.getContext());
5276 patterns.add<ResolvePodReadBackedArrayReadOp>(patterns.getContext(), *this);
5277 patterns.add<ResolvePodReadBackedArrayLengthOp>(patterns.getContext(), *this);
5278 patterns.add<ResolveDeferredSplitPodArrayCastOp>(patterns.getContext(), *this);
5279 patterns.add<ResolveVirtualPodWriteOp>(patterns.getContext(), *this);
5280 patterns.add<ResolveVirtualPodReadOp>(patterns.getContext(), *this);
5281}
5282
5283void Step3Resolver::addLateResolutionPatterns(RewritePatternSet &patterns) {
5284 patterns.add<ResolvePodReadBackedArrayReadOp>(patterns.getContext(), *this);
5285 patterns.add<ResolvePodReadBackedArrayLengthOp>(patterns.getContext(), *this);
5286 patterns.add<ResolveDeferredSplitPodArrayCastOp>(patterns.getContext(), *this);
5287 patterns.add<ResolveVirtualPodWriteOp>(patterns.getContext(), *this);
5288 patterns.add<ResolveVirtualPodReadOp>(patterns.getContext(), *this);
5289}
5290
5291void Step3Resolver::addPostConversionPatterns(RewritePatternSet &patterns) {
5292 patterns.add<SplitVirtualPodInEmitEqualityPattern>(patterns.getContext(), *this);
5293}
5294
5295void Step3Resolver::configureLateVirtualPodLegality(ConversionTarget &target) const {
5296 target.addDynamicallyLegalOp<WritePodOp>([this](WritePodOp op) {
5297 return !lookupVirtualPodLeafMap(op.getPodRef(), virtualPods) ||
5298 isInsideSupportedScfRegion(op.getOperation());
5299 });
5300 target.addDynamicallyLegalOp<ArrayLengthOp>([](ArrayLengthOp op) {
5301 return RejectRaggedNestedLeafArrayLengthOp::legal(op) &&
5302 ResolvePodReadBackedArrayLengthOp::legal(op);
5303 });
5304 target.addDynamicallyLegalOp<ReadArrayOp>([this](ReadArrayOp op) {
5305 return !ResolvePodReadBackedArrayReadOp::canResolve(op, *this);
5306 });
5307 target.addDynamicallyLegalOp<UnrealizedConversionCastOp>([this](UnrealizedConversionCastOp op) {
5308 return !ResolveDeferredSplitPodArrayCastOp::canResolve(op, *this);
5309 });
5310 target.addDynamicallyLegalOp<ReadPodOp>([this](ReadPodOp op) {
5311 return !canResolveVirtualPodRead(op, virtualPods);
5312 });
5313}
5314
5315bool Step3Resolver::hasResolvableLateVirtualPodOps(ModuleOp modOp) const {
5316 return walkContains<Operation *>(*modOp, [this](Operation *op) {
5317 return TypeSwitch<Operation *, bool>(op)
5318 .Case<WritePodOp>([this](auto writeOp) {
5319 return lookupVirtualPodLeafMap(writeOp.getPodRef(), virtualPods) &&
5320 !isInsideSupportedScfRegion(writeOp.getOperation());
5321 })
5322 .Case<ReadPodOp>([this](auto readOp) {
5323 return canResolveVirtualPodRead(readOp, virtualPods);
5324 })
5325 .Case<ReadArrayOp>([this](auto readOp) {
5326 return ResolvePodReadBackedArrayReadOp::canResolve(readOp, *this);
5327 })
5328 .Case<ArrayLengthOp>([](auto lenOp) {
5329 return !ResolvePodReadBackedArrayLengthOp::legal(lenOp);
5330 })
5331 .Case<UnrealizedConversionCastOp>([this](auto castOp) {
5332 return ResolveDeferredSplitPodArrayCastOp::canResolve(castOp, *this);
5333 }).Default([](Operation *) { return false; });
5334 });
5335}
5336
5337void Step3Resolver::materializeRemainingVirtualPods(ModuleOp modOp) {
5338 OpBuilder builder(modOp.getContext());
5339 modOp.walk([this, &builder](NewPodOp newPod) {
5340 auto it = virtualPods.find(newPod.getResult());
5341 if (it == virtualPods.end() || newPod.use_empty()) {
5342 return;
5343 }
5344 builder.setInsertionPointAfter(findVirtualPodMaterializationAnchor(newPod, it->second));
5345 materializeVirtualPod(builder, newPod, it->second);
5346 });
5347}
5348
5351static LogicalResult
5352step3(ModuleOp modOp, SymbolTableCollection &symTables, const MemberReplacementMap &memberRepMap) {
5353 MLIRContext *ctx = modOp.getContext();
5354 Step3Resolver resolver;
5355 resolver.rehydrateVirtualPodPlaceholders(modOp);
5356
5357 RewritePatternSet preConversionPatterns(ctx);
5358 resolver.addPreConversionPatterns(preConversionPatterns);
5359 if (failed(applyPatternsGreedily(
5360 modOp->getRegion(0), std::move(preConversionPatterns),
5361 GreedyRewriteConfig {.fold = false, .cseConstants = false}
5362 ))) {
5363 return failure();
5364 }
5365
5366 RewritePatternSet patterns(ctx);
5367 resolver.addConversionPatterns(patterns, symTables, memberRepMap);
5368
5369 ConversionTarget target(*ctx);
5370 baseTargetSetup(target);
5371 target.addDynamicallyLegalOp<NewPodOp>(SplitInitFromNewPodOp::legal);
5372 target.addDynamicallyLegalOp<CreateArrayOp>(SplitPodElementCreateArrayOp::legal);
5373 target.addDynamicallyLegalOp<FuncDefOp>(SplitPodInFuncDefOp::legal);
5374 target.addDynamicallyLegalOp<ReturnOp>(SplitPodInReturnOp::legal);
5375 target.addDynamicallyLegalOp<CallOp>(SplitPodInCallOp::legal);
5376 target.addDynamicallyLegalOp<MemberWriteOp>(SplitPodInMemberWriteOp::legal);
5377 target.addDynamicallyLegalOp<MemberReadOp>(SplitPodInMemberReadOp::legal);
5378 resolver.configureLateVirtualPodLegality(target);
5379
5380 LLVM_DEBUG(llvm::dbgs() << "Begin step 3: update/split other pod ops\n";);
5381 if (failed(applyFullConversion(modOp, target, std::move(patterns)))) {
5382 return failure();
5383 }
5384
5385 // Step 3's full conversion can expose a second wave of virtual-POD cleanup
5386 // opportunities. For example, resolving one deferred array bridge can make a
5387 // previously-illegal pod.read/pod.write or cast finally resolvable. Scan for
5388 // exactly those ops here so we can drive a small fixpoint loop before the
5389 // final materialization stage.
5390 // Rebuild the late-resolution target and pattern set each round because the
5391 // set of remaining virtual-POD placeholders shrinks as patterns fire. This
5392 // phase intentionally uses partial conversion: each round resolves whatever
5393 // is now legal to lower, then we rescan for any newly-exposed opportunities.
5394 auto runLateResolutionRound = [&]() {
5395 RewritePatternSet lateResolutionPatterns(ctx);
5396 resolver.addLateResolutionPatterns(lateResolutionPatterns);
5397
5398 ConversionTarget lateResolutionTarget(*ctx);
5399 baseTargetSetup(lateResolutionTarget);
5400 lateResolutionTarget.addLegalOp<ModuleOp>();
5401 resolver.configureLateVirtualPodLegality(lateResolutionTarget);
5402
5403 return applyPartialConversion(modOp, lateResolutionTarget, std::move(lateResolutionPatterns));
5404 };
5405
5406 // Iterate to a fixpoint before post-processing materializes any surviving
5407 // virtual PODs. A bounded loop guards against accidental non-progress.
5408 for (unsigned lateResolutionRounds = 0;
5409 lateResolutionRounds < 64 && resolver.hasResolvableLateVirtualPodOps(modOp);
5410 ++lateResolutionRounds) {
5411 if (failed(runLateResolutionRound())) {
5412 return failure();
5413 }
5414 }
5415 if (resolver.hasResolvableLateVirtualPodOps(modOp)) {
5416 modOp.emitError("late virtual POD resolution did not reach a fixpoint");
5417 return failure();
5418 }
5419
5420 RewritePatternSet postConversionPatterns(ctx);
5421 resolver.addPostConversionPatterns(postConversionPatterns);
5422 if (failed(applyPatternsGreedily(
5423 modOp->getRegion(0), std::move(postConversionPatterns),
5424 GreedyRewriteConfig {.fold = false, .cseConstants = false}
5425 ))) {
5426 return failure();
5427 }
5428 if (failed(promoteFunctionArgCastsToSignature(modOp, symTables))) {
5429 return failure();
5430 }
5431
5432 resolver.materializeRemainingVirtualPods(modOp);
5433
5434 bool erasedDeadPlaceholderOps = false;
5435 do {
5436 SmallVector<Operation *> deadPlaceholderOps;
5437 modOp->walk<WalkOrder::PostOrder>([&deadPlaceholderOps](Operation *op) {
5438 if (auto readOp = llvm::dyn_cast<ReadPodOp>(op)) {
5439 if (readOp.getResult().use_empty()) {
5440 deadPlaceholderOps.push_back(op);
5441 }
5442 return;
5443 }
5444
5445 if (auto castOp = llvm::dyn_cast<UnrealizedConversionCastOp>(op)) {
5446 if (llvm::all_of(castOp.getResults(), [](Value result) { return result.use_empty(); })) {
5447 deadPlaceholderOps.push_back(op);
5448 }
5449 }
5450 });
5451 for (Operation *op : deadPlaceholderOps) {
5452 op->erase();
5453 }
5454 erasedDeadPlaceholderOps = !deadPlaceholderOps.empty();
5455 } while (erasedDeadPlaceholderOps);
5456
5457 SmallVector<Operation *> deadOps;
5458 modOp->walk<WalkOrder::PostOrder>([&](Operation *op) {
5459 if (op != modOp.getOperation() && isOpTriviallyDead(op)) {
5460 deadOps.push_back(op);
5461 }
5462 });
5463 for (Operation *op : deadOps) {
5464 op->erase();
5465 }
5466 return success();
5467}
5468
5473static bool isValueDefinedInside(Operation *ancestor, Value value) {
5474 if (Operation *defOp = value.getDefiningOp()) {
5475 return ancestor->isAncestor(defOp);
5476 }
5477
5478 auto blockArg = llvm::dyn_cast<BlockArgument>(value);
5479 Operation *parentOp = blockArg.getOwner()->getParentOp();
5480 return parentOp && ancestor->isAncestor(parentOp);
5481}
5482
5484static WritePodOp findPrecedingWriteForIfRead(ReadPodOp readOp) {
5485 auto ifOp = readOp->getParentOfType<scf::IfOp>();
5486 if (!ifOp || readOp->getBlock()->getParentOp() != ifOp.getOperation()) {
5487 return nullptr;
5488 }
5489 if (hasEarlierWriteInBlock(readOp)) {
5490 return nullptr;
5491 }
5492
5493 Block *ifBlock = ifOp->getBlock();
5494 if (!ifBlock) {
5495 return nullptr;
5496 }
5497
5498 Value podRef = readOp.getPodRef();
5499 StringAttr recordName = readOp.getRecordNameAttr();
5500 WritePodOp replacement = nullptr;
5501 for (Operation &op : *ifBlock) {
5502 if (&op == ifOp.getOperation()) {
5503 break;
5504 }
5505
5506 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5507 if (isSamePodRecord(writeOp, podRef, recordName)) {
5508 replacement = writeOp;
5509 }
5510 continue;
5511 }
5512
5513 if (hasNestedWriteToRecord(op, podRef, recordName)) {
5514 replacement = nullptr;
5515 }
5516 }
5517
5518 return replacement;
5519}
5520
5522class FoldReadAfterWriteInBlockPattern final : public OpRewritePattern<ReadPodOp> {
5523public:
5524 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5525
5526 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5527 if (WritePodOp writeOp = findNearestForwardableWriteInBlock(readOp)) {
5528 rewriter.replaceOp(readOp, writeOp.getValue());
5529 return success();
5530 }
5531 return failure();
5532 }
5533};
5534
5536class ReplaceIfReadPattern final : public OpRewritePattern<ReadPodOp> {
5537public:
5538 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5539
5540 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5541 auto ifOp = readOp->getParentOfType<scf::IfOp>();
5542 if (!ifOp || readOp->getBlock()->getParentOp() != ifOp.getOperation()) {
5543 return failure();
5544 }
5545 if (isValueDefinedInside(ifOp, readOp.getPodRef()) || hasEarlierWriteInBlock(readOp)) {
5546 return failure();
5547 }
5548
5549 if (WritePodOp writeOp = findPrecedingWriteForIfRead(readOp)) {
5550 rewriter.replaceOp(readOp, writeOp.getValue());
5551 return success();
5552 }
5553
5554 rewriter.setInsertionPoint(ifOp);
5555 rewriter.replaceOp(
5556 readOp, genRead(rewriter, readOp.getLoc(), readOp.getPodRef(), readOp.getRecordNameAttr())
5557 .getResult()
5558 );
5559 return success();
5560 }
5561};
5562
5576class FoldIfCarriedPodReadAfterWritePattern final : public OpRewritePattern<ReadPodOp> {
5577public:
5578 using OpRewritePattern<ReadPodOp>::OpRewritePattern;
5579
5580 LogicalResult matchAndRewrite(ReadPodOp readOp, PatternRewriter &rewriter) const override {
5581 auto podRes = dyn_cast<OpResult>(readOp.getPodRef());
5582 if (!podRes) {
5583 return failure();
5584 }
5585
5586 auto ifOp = dyn_cast<scf::IfOp>(podRes.getOwner());
5587 if (!ifOp) {
5588 return failure();
5589 }
5590
5591 auto writeOp = dyn_cast_or_null<WritePodOp>(readOp->getPrevNode());
5592 if (!writeOp || writeOp.getRecordNameAttr() != readOp.getRecordNameAttr()) {
5593 return failure();
5594 }
5595
5596 auto valueRes = dyn_cast<OpResult>(writeOp.getValue());
5597 if (!valueRes || valueRes.getOwner() != ifOp.getOperation()) {
5598 return failure();
5599 }
5600
5601 Value carriedPod = writeOp.getPodRef();
5602 unsigned podResultIndex = podRes.getResultNumber();
5603
5604 auto thenYield = dyn_cast<scf::YieldOp>(ifOp.thenBlock()->getTerminator());
5605 if (!thenYield || thenYield.getOperand(podResultIndex) != carriedPod) {
5606 return failure();
5607 }
5608
5609 Region &elseRegion = ifOp.getElseRegion();
5610 if (Block *elseBlock = elseRegion.empty() ? nullptr : &elseRegion.front()) {
5611 auto elseYield = dyn_cast<scf::YieldOp>(elseBlock->getTerminator());
5612 if (!elseYield || elseYield.getOperand(podResultIndex) != carriedPod) {
5613 return failure();
5614 }
5615 }
5616
5617 rewriter.replaceOp(readOp, valueRes);
5618 return success();
5619 }
5620};
5621
5626struct IfWriteSlot {
5627 Value podRef;
5628 StringAttr recordName;
5629 Type type;
5630 WritePodOp thenWrite;
5631 WritePodOp elseWrite;
5632 Value incomingValue;
5633};
5634
5636static IfWriteSlot *
5637lookupSlot(SmallVectorImpl<IfWriteSlot> &slots, Value podRef, StringAttr recordName) {
5638 for (IfWriteSlot &slot : slots) {
5639 if (slot.podRef == podRef && slot.recordName == recordName) {
5640 return &slot;
5641 }
5642 }
5643 return nullptr;
5644}
5645
5647static IfWriteSlot &getOrCreateSlot(
5648 SmallVectorImpl<IfWriteSlot> &slots, Value podRef, StringAttr recordName, Type type
5649) {
5650 if (IfWriteSlot *slot = lookupSlot(slots, podRef, recordName)) {
5651 return *slot;
5652 }
5653 slots.push_back(IfWriteSlot {podRef, recordName, type, nullptr, nullptr, Value()});
5654 return slots.back();
5655}
5656
5658static Block *getElseBlockOrNull(scf::IfOp ifOp) {
5659 return ifOp.getElseRegion().empty() ? nullptr : &ifOp.getElseRegion().front();
5660}
5661
5663static void
5664collectDirectWrites(Block *block, bool isThenBlock, SmallVectorImpl<IfWriteSlot> &slots) {
5665 if (!block) {
5666 return;
5667 }
5668
5669 for (Operation &op : *block) {
5670 if (op.hasTrait<OpTrait::IsTerminator>()) {
5671 break;
5672 }
5673
5674 auto writeOp = dyn_cast<WritePodOp>(&op);
5675 if (!writeOp) {
5676 continue;
5677 }
5678
5679 IfWriteSlot &slot = getOrCreateSlot(
5680 slots, writeOp.getPodRef(), writeOp.getRecordNameAttr(), writeOp.getValue().getType()
5681 );
5682 if (isThenBlock) {
5683 slot.thenWrite = writeOp;
5684 } else {
5685 slot.elseWrite = writeOp;
5686 }
5687 }
5688}
5689
5694static bool branchSlotCanBeLifted(Block *block, Value podRef, StringAttr recordName) {
5695 if (!block) {
5696 return true;
5697 }
5698
5699 bool seenDirectWrite = false;
5700 for (Operation &op : *block) {
5701 if (op.hasTrait<OpTrait::IsTerminator>()) {
5702 return true;
5703 }
5704
5705 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5706 if (isSamePodRecord(writeOp, podRef, recordName)) {
5707 seenDirectWrite = true;
5708 continue;
5709 }
5710 }
5711
5712 if (hasNestedWriteToRecord(op, podRef, recordName)) {
5713 return false;
5714 }
5715 if (seenDirectWrite && (hasReadFromRecord(op, podRef, recordName) || hasValueUse(op, podRef))) {
5716 return false;
5717 }
5718 }
5719 return true;
5720}
5721
5723static bool isLiftedWrite(Operation &op, ArrayRef<IfWriteSlot> slots) {
5724 auto writeOp = dyn_cast<WritePodOp>(&op);
5725 return writeOp && llvm::any_of(slots, [&writeOp](const IfWriteSlot &slot) {
5726 return isSamePodRecord(writeOp, slot.podRef, slot.recordName);
5727 });
5728}
5729
5731static scf::YieldOp getYieldOp(Block &block) {
5732 auto yieldOp = dyn_cast<scf::YieldOp>(block.getTerminator());
5733 assert(yieldOp && "expected scf.if branch to terminate with scf.yield");
5734 return yieldOp;
5735}
5736
5738static void dropTerminatorIfPresent(Block &block) {
5739 if (!block.empty() && block.back().hasTrait<OpTrait::IsTerminator>()) {
5740 block.back().erase();
5741 }
5742}
5743
5745static void
5746moveBranchWithoutLiftedWrites(Block *srcBlock, Block &destBlock, ArrayRef<IfWriteSlot> slots) {
5747 if (srcBlock) {
5748 for (auto it = srcBlock->begin(), end = srcBlock->end(); it != end;) {
5749 Operation &op = *it++;
5750 if (op.hasTrait<OpTrait::IsTerminator>() || isLiftedWrite(op, slots)) {
5751 continue;
5752 }
5753 op.moveBefore(&destBlock, destBlock.end());
5754 }
5755 }
5756}
5757
5760static void appendYield(
5761 OpBuilder &bldr, Location loc, Block &block, ValueRange priorYieldValues,
5762 ArrayRef<IfWriteSlot> slots, bool isThenBlock, scf::YieldOp originalYield = nullptr
5763) {
5764 SmallVector<Value> yieldValues = llvm::to_vector(priorYieldValues);
5765 llvm::append_range(yieldValues, llvm::map_range(slots, [isThenBlock](const IfWriteSlot &slot) {
5766 WritePodOp writeOp = isThenBlock ? slot.thenWrite : slot.elseWrite;
5767 return writeOp ? writeOp.getValue() : slot.incomingValue;
5768 }));
5769
5770 bldr.setInsertionPointToEnd(&block);
5771 auto newYield = bldr.create<scf::YieldOp>(loc, yieldValues);
5772 if (originalYield) {
5773 preserveDiscardableAttrs(originalYield, newYield);
5774 }
5775}
5776
5782struct LoopPodSlot {
5783 Value podRef;
5784 StringAttr recordName;
5785 Type type;
5786
5788 bool matches(Value findPodRef, StringAttr findRecordName) const {
5789 return this->podRef == findPodRef && this->recordName == findRecordName;
5790 }
5791};
5792
5794static LoopPodSlot *
5795lookupLoopSlot(SmallVectorImpl<LoopPodSlot> &slots, Value podRef, StringAttr recordName) {
5796 auto *it = llvm::find_if(slots, [&podRef, &recordName](const LoopPodSlot &slot) {
5797 return slot.matches(podRef, recordName);
5798 });
5799 return it == slots.end() ? nullptr : &*it;
5800}
5801
5803static bool hasLoopSlot(ArrayRef<LoopPodSlot> slots, Value podRef, StringAttr recordName) {
5804 const auto *it = llvm::find_if(slots, [&podRef, &recordName](const LoopPodSlot &slot) {
5805 return slot.matches(podRef, recordName);
5806 });
5807 return it != slots.end();
5808}
5809
5811static LoopPodSlot &getOrCreateLoopSlot(
5812 SmallVectorImpl<LoopPodSlot> &slots, Value podRef, StringAttr recordName, Type type
5813) {
5814 if (LoopPodSlot *slot = lookupLoopSlot(slots, podRef, recordName)) {
5815 return *slot;
5816 }
5817 slots.push_back(LoopPodSlot {podRef, recordName, type});
5818 return slots.back();
5819}
5820
5822static std::optional<size_t>
5823findLoopSlotIndex(ArrayRef<LoopPodSlot> slots, Value podRef, StringAttr recordName) {
5824 for (auto [idx, slot] : llvm::enumerate(slots)) {
5825 if (slot.podRef == podRef && slot.recordName == recordName) {
5826 return idx;
5827 }
5828 }
5829 return std::nullopt;
5830}
5831
5834static void
5835collectDirectLoopPodSlots(Block &block, Operation *ancestor, SmallVectorImpl<LoopPodSlot> &slots) {
5836 for (Operation &op : block) {
5837 if (auto readOp = dyn_cast<ReadPodOp>(&op)) {
5838 if (!isValueDefinedInside(ancestor, readOp.getPodRef())) {
5839 getOrCreateLoopSlot(
5840 slots, readOp.getPodRef(), readOp.getRecordNameAttr(), readOp.getType()
5841 );
5842 }
5843 continue;
5844 }
5845
5846 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5847 if (!isValueDefinedInside(ancestor, writeOp.getPodRef())) {
5848 getOrCreateLoopSlot(
5849 slots, writeOp.getPodRef(), writeOp.getRecordNameAttr(), writeOp.getValue().getType()
5850 );
5851 }
5852 }
5853 }
5854}
5855
5857static bool opUsesTrackedPodRefDirectly(Operation &op, ArrayRef<LoopPodSlot> slots) {
5858 return llvm::any_of(op.getOperands(), [&slots](Value operand) {
5859 return llvm::any_of(slots, [&operand](const LoopPodSlot &slot) {
5860 return slot.podRef == operand;
5861 });
5862 });
5863}
5864
5866static bool hasNestedTrackedPodAccess(Operation &op, ArrayRef<LoopPodSlot> slots) {
5867 return op
5868 .walk([&op, &slots](Operation *nestedOp) {
5869 if (nestedOp == &op) {
5870 return WalkResult::advance();
5871 }
5872
5873 if (auto readOp = dyn_cast<ReadPodOp>(nestedOp)) {
5874 if (hasLoopSlot(slots, readOp.getPodRef(), readOp.getRecordNameAttr())) {
5875 return WalkResult::interrupt();
5876 }
5877 return WalkResult::advance();
5878 }
5879
5880 if (auto writeOp = dyn_cast<WritePodOp>(nestedOp)) {
5881 if (hasLoopSlot(slots, writeOp.getPodRef(), writeOp.getRecordNameAttr())) {
5882 return WalkResult::interrupt();
5883 }
5884 }
5885 return WalkResult::advance();
5886 }).wasInterrupted();
5887}
5888
5891static bool hasUnliftableLoopPodUses(Block &block, ArrayRef<LoopPodSlot> slots) {
5892 for (Operation &op : block) {
5893 if (isa<ReadPodOp, WritePodOp>(op)) {
5894 continue;
5895 }
5896 if (opUsesTrackedPodRefDirectly(op, slots) || hasNestedTrackedPodAccess(op, slots)) {
5897 return true;
5898 }
5899 }
5900 return false;
5901}
5902
5904template <typename RewriteTerminatorFn>
5905static void cloneLoopBodyWithLiftedPodSlots(
5906 Block &source, PatternRewriter &rewriter, IRMapping &mapping, ArrayRef<LoopPodSlot> slots,
5907 SmallVectorImpl<Value> &slotValues, RewriteTerminatorFn &&rewriteTerminator
5908) {
5909 for (Operation &op : source) {
5910 if (rewriteTerminator(op)) {
5911 continue;
5912 }
5913
5914 if (auto readOp = dyn_cast<ReadPodOp>(&op)) {
5915 if (std::optional<size_t> slotIdx =
5916 findLoopSlotIndex(slots, readOp.getPodRef(), readOp.getRecordNameAttr())) {
5917 mapping.map(readOp.getResult(), slotValues[*slotIdx]);
5918 continue;
5919 }
5920 }
5921
5922 if (auto writeOp = dyn_cast<WritePodOp>(&op)) {
5923 if (std::optional<size_t> slotIdx =
5924 findLoopSlotIndex(slots, writeOp.getPodRef(), writeOp.getRecordNameAttr())) {
5925 slotValues[*slotIdx] = mapping.lookupOrDefault(writeOp.getValue());
5926 continue;
5927 }
5928 }
5929
5930 rewriter.clone(op, mapping);
5931 }
5932}
5933
5935static void appendIncomingLoopSlotValues(
5936 PatternRewriter &rewriter, Location loc, ArrayRef<LoopPodSlot> slots,
5937 SmallVectorImpl<Value> &values, SmallVectorImpl<Type> *resultTypes = nullptr
5938) {
5939 for (const LoopPodSlot &slot : slots) {
5940 values.push_back(genRead(rewriter, loc, slot.podRef, slot.recordName).getResult());
5941 if (resultTypes) {
5942 resultTypes->push_back(slot.type);
5943 }
5944 }
5945}
5946
5948template <typename ValueRangeLike>
5949static SmallVector<Value>
5950collectTrailingLoopSlotValues(ValueRangeLike values, size_t base, size_t slotCount) {
5951 SmallVector<Value> slotValues;
5952 slotValues.reserve(slotCount);
5953 for (size_t idx = 0; idx < slotCount; ++idx) {
5954 slotValues.push_back(values[llzk::checkedCast<unsigned>(base + idx)]);
5955 }
5956 return slotValues;
5957}
5958
5960static SmallVector<Value>
5961remapValuesAndAppendLoopSlots(ValueRange values, IRMapping &mapping, ValueRange slotValues) {
5962 SmallVector<Value> remappedValues = llvm::map_to_vector(values, [&mapping](Value value) {
5963 return mapping.lookupOrDefault(value);
5964 });
5965 llvm::append_range(remappedValues, slotValues);
5966 return remappedValues;
5967}
5968
5970template <typename GetResultFn>
5971static void writeBackLoopSlotResults(
5972 PatternRewriter &rewriter, Location loc, ArrayRef<LoopPodSlot> slots,
5973 unsigned originalResultCount, GetResultFn &&getResult
5974) {
5975 for (auto [idx, slot] : llvm::enumerate(slots)) {
5976 genWrite(
5977 rewriter, loc, slot.podRef, slot.recordName,
5978 getResult(originalResultCount + llzk::checkedCast<unsigned>(idx))
5979 );
5980 }
5981}
5982
5986class LiftPodWritesFromIfBlocksPattern final : public OpRewritePattern<scf::IfOp> {
5987public:
5988 using OpRewritePattern<scf::IfOp>::OpRewritePattern;
5989
5990 LogicalResult matchAndRewrite(scf::IfOp ifOp, PatternRewriter &rewriter) const override {
5991 SmallVector<IfWriteSlot> slots;
5992 Block &thenBlock = *ifOp.thenBlock();
5993 Block *elseBlock = getElseBlockOrNull(ifOp);
5994 collectDirectWrites(&thenBlock, true, slots);
5995 collectDirectWrites(elseBlock, false, slots);
5996 if (slots.empty()) {
5997 return failure();
5998 }
5999
6000 llvm::erase_if(slots, [&](const IfWriteSlot &slot) {
6001 return isValueDefinedInside(ifOp, slot.podRef) ||
6002 !branchSlotCanBeLifted(&thenBlock, slot.podRef, slot.recordName) ||
6003 !branchSlotCanBeLifted(elseBlock, slot.podRef, slot.recordName);
6004 });
6005 if (slots.empty()) {
6006 return failure();
6007 }
6008
6009 for (IfWriteSlot &slot : slots) {
6010 if (slot.thenWrite && slot.elseWrite) {
6011 continue;
6012 }
6013 rewriter.setInsertionPoint(ifOp);
6014 slot.incomingValue =
6015 genRead(rewriter, ifOp.getLoc(), slot.podRef, slot.recordName).getResult();
6016 }
6017
6018 SmallVector<Type> resultTypes = llvm::to_vector(ifOp.getResultTypes());
6019 llvm::append_range(resultTypes, llvm::map_range(slots, [](auto slot) { return slot.type; }));
6020
6021 scf::YieldOp thenYieldOp = getYieldOp(thenBlock);
6022 SmallVector<Value> originalThenYields;
6023 if (!ifOp.getResults().empty()) {
6024 originalThenYields.append(thenYieldOp.getOperands().begin(), thenYieldOp.getOperands().end());
6025 }
6026
6027 scf::YieldOp elseYieldOp = elseBlock ? getYieldOp(*elseBlock) : nullptr;
6028 SmallVector<Value> originalElseYields;
6029 if (elseBlock && !ifOp.getResults().empty()) {
6030 originalElseYields.append(elseYieldOp.getOperands().begin(), elseYieldOp.getOperands().end());
6031 }
6032
6033 rewriter.setInsertionPoint(ifOp);
6034 auto newIf = rewriter.create<scf::IfOp>(ifOp.getLoc(), resultTypes, ifOp.getCondition(), true);
6035 Block &newThenBlock = *newIf.thenBlock();
6036 Block &newElseBlock = *newIf.elseBlock();
6037 dropTerminatorIfPresent(newThenBlock);
6038 dropTerminatorIfPresent(newElseBlock);
6039
6040 moveBranchWithoutLiftedWrites(&thenBlock, newThenBlock, slots);
6041 moveBranchWithoutLiftedWrites(elseBlock, newElseBlock, slots);
6042 appendYield(
6043 rewriter, ifOp.getLoc(), newThenBlock, originalThenYields, slots, true, thenYieldOp
6044 );
6045 appendYield(
6046 rewriter, ifOp.getLoc(), newElseBlock, originalElseYields, slots, false, elseYieldOp
6047 );
6048
6049 rewriter.setInsertionPointAfter(newIf);
6050 unsigned originalResultCount = ifOp.getNumResults();
6051 for (auto [idx, slot] : llvm::enumerate(slots)) {
6052 genWrite(
6053 rewriter, ifOp.getLoc(), slot.podRef, slot.recordName,
6054 newIf.getResult(originalResultCount + idx)
6055 );
6056 }
6057
6058 rewriter.replaceOp(ifOp, newIf.getResults().take_front(originalResultCount));
6059 return success();
6060 }
6061};
6062
6065class LiftPodAccessesFromForLoopPattern final : public OpRewritePattern<scf::ForOp> {
6066public:
6067 using OpRewritePattern<scf::ForOp>::OpRewritePattern;
6068
6069 LogicalResult matchAndRewrite(scf::ForOp forOp, PatternRewriter &rewriter) const override {
6070 Block &body = *forOp.getBody();
6071 SmallVector<LoopPodSlot> slots;
6072 collectDirectLoopPodSlots(body, forOp.getOperation(), slots);
6073 if (slots.empty() || hasUnliftableLoopPodUses(body, slots)) {
6074 return failure();
6075 }
6076
6077 Location loc = forOp.getLoc();
6078
6079 SmallVector<Value> newInitArgs = llvm::to_vector(forOp.getInitArgs());
6080 rewriter.setInsertionPoint(forOp);
6081 appendIncomingLoopSlotValues(rewriter, loc, slots, newInitArgs);
6082
6083 auto newFor = rewriter.create<scf::ForOp>(
6084 loc, forOp.getLowerBound(), forOp.getUpperBound(), forOp.getStep(), newInitArgs
6085 );
6086 newFor->setAttrs(forOp->getAttrs());
6087
6088 Block &newBody = *newFor.getBody();
6089 dropTerminatorIfPresent(newBody);
6090
6091 IRMapping mapping;
6092 mapping.map(forOp.getInductionVar(), newFor.getInductionVar());
6093 for (auto [idx, oldArg] : llvm::enumerate(forOp.getRegionIterArgs())) {
6094 mapping.map(oldArg, newFor.getRegionIterArg(idx));
6095 }
6096
6097 SmallVector<Value> slotValues = collectTrailingLoopSlotValues(
6098 newFor.getRegionIterArgs(), forOp.getNumRegionIterArgs(), slots.size()
6099 );
6100
6101 rewriter.setInsertionPointToEnd(&newBody);
6102 cloneLoopBodyWithLiftedPodSlots(body, rewriter, mapping, slots, slotValues, [&](Operation &op) {
6103 if (auto yieldOp = dyn_cast<scf::YieldOp>(&op)) {
6104 SmallVector<Value> yieldValues =
6105 remapValuesAndAppendLoopSlots(yieldOp.getOperands(), mapping, slotValues);
6106 preserveDiscardableAttrs(
6107 yieldOp, rewriter.create<scf::YieldOp>(yieldOp.getLoc(), yieldValues)
6108 );
6109 return true;
6110 }
6111 return false;
6112 });
6113
6114 rewriter.setInsertionPointAfter(newFor);
6115 writeBackLoopSlotResults(
6116 rewriter, loc, slots, forOp.getNumResults(),
6117 [&newFor](unsigned resultIdx) { return newFor.getResult(resultIdx); }
6118 );
6119
6120 rewriter.replaceOp(forOp, newFor.getResults().take_front(forOp.getNumResults()));
6121 return success();
6122 }
6123};
6124
6127class LiftPodAccessesFromWhileLoopPattern final : public OpRewritePattern<scf::WhileOp> {
6128public:
6129 using OpRewritePattern<scf::WhileOp>::OpRewritePattern;
6130
6131 LogicalResult matchAndRewrite(scf::WhileOp whileOp, PatternRewriter &rewriter) const override {
6132 Block &beforeBody = *whileOp.getBeforeBody();
6133 Block &afterBody = *whileOp.getAfterBody();
6134
6135 SmallVector<LoopPodSlot> slots;
6136 collectDirectLoopPodSlots(beforeBody, whileOp.getOperation(), slots);
6137 collectDirectLoopPodSlots(afterBody, whileOp.getOperation(), slots);
6138 if (slots.empty() || hasUnliftableLoopPodUses(beforeBody, slots) ||
6139 hasUnliftableLoopPodUses(afterBody, slots)) {
6140 return failure();
6141 }
6142
6143 Location loc = whileOp.getLoc();
6144
6145 SmallVector<Value> newInits = llvm::to_vector(whileOp.getInits());
6146 SmallVector<Type> newResultTypes = llvm::to_vector(whileOp.getResultTypes());
6147 rewriter.setInsertionPoint(whileOp);
6148 appendIncomingLoopSlotValues(rewriter, loc, slots, newInits, &newResultTypes);
6149
6150 auto newWhile = rewriter.create<scf::WhileOp>(loc, newResultTypes, newInits, nullptr, nullptr);
6151 newWhile->setAttrs(whileOp->getAttrs());
6152
6153 Block &newBeforeBody = *newWhile.getBeforeBody();
6154 Block &newAfterBody = *newWhile.getAfterBody();
6155 dropTerminatorIfPresent(newBeforeBody);
6156 dropTerminatorIfPresent(newAfterBody);
6157
6158 IRMapping beforeMapping;
6159 for (auto [oldArg, newArg] : llvm::zip_equal(
6160 whileOp.getBeforeArguments(),
6161 newWhile.getBeforeArguments().take_front(whileOp.getBeforeArguments().size())
6162 )) {
6163 beforeMapping.map(oldArg, newArg);
6164 }
6165
6166 SmallVector<Value> beforeSlotValues = collectTrailingLoopSlotValues(
6167 newWhile.getBeforeArguments(), whileOp.getBeforeArguments().size(), slots.size()
6168 );
6169
6170 rewriter.setInsertionPointToEnd(&newBeforeBody);
6171 cloneLoopBodyWithLiftedPodSlots(
6172 beforeBody, rewriter, beforeMapping, slots, beforeSlotValues, [&](Operation &op) {
6173 if (auto conditionOp = dyn_cast<scf::ConditionOp>(&op)) {
6174 SmallVector<Value> conditionArgs =
6175 remapValuesAndAppendLoopSlots(conditionOp.getArgs(), beforeMapping, beforeSlotValues);
6176 preserveDiscardableAttrs(
6177 conditionOp,
6178 rewriter.create<scf::ConditionOp>(
6179 conditionOp.getLoc(), beforeMapping.lookupOrDefault(conditionOp.getCondition()),
6180 conditionArgs
6181 )
6182 );
6183 return true;
6184 }
6185 return false;
6186 }
6187 );
6188
6189 IRMapping afterMapping;
6190 for (auto [oldArg, newArg] : llvm::zip_equal(
6191 whileOp.getAfterArguments(),
6192 newWhile.getAfterArguments().take_front(whileOp.getAfterArguments().size())
6193 )) {
6194 afterMapping.map(oldArg, newArg);
6195 }
6196
6197 SmallVector<Value> afterSlotValues = collectTrailingLoopSlotValues(
6198 newWhile.getAfterArguments(), whileOp.getAfterArguments().size(), slots.size()
6199 );
6200
6201 rewriter.setInsertionPointToEnd(&newAfterBody);
6202 cloneLoopBodyWithLiftedPodSlots(
6203 afterBody, rewriter, afterMapping, slots, afterSlotValues, [&](Operation &op) {
6204 if (auto yieldOp = dyn_cast<scf::YieldOp>(&op)) {
6205 SmallVector<Value> yieldValues =
6206 remapValuesAndAppendLoopSlots(yieldOp.getOperands(), afterMapping, afterSlotValues);
6207 preserveDiscardableAttrs(
6208 yieldOp, rewriter.create<scf::YieldOp>(yieldOp.getLoc(), yieldValues)
6209 );
6210 return true;
6211 }
6212 return false;
6213 }
6214 );
6215
6216 rewriter.setInsertionPointAfter(newWhile);
6217 writeBackLoopSlotResults(
6218 rewriter, loc, slots, whileOp.getNumResults(),
6219 [&newWhile](unsigned resultIdx) { return newWhile.getResult(resultIdx); }
6220 );
6221
6222 rewriter.replaceOp(whileOp, newWhile.getResults().take_front(whileOp.getNumResults()));
6223 return success();
6224 }
6225};
6226
6232class SplitPodInEmitEqualityPattern final : public OpRewritePattern<constrain::EmitEqualityOp> {
6233 CompatiblePodLeafMaterializationMap &materializedLeaves;
6234
6235public:
6236 SplitPodInEmitEqualityPattern(
6237 MLIRContext *ctx, CompatiblePodLeafMaterializationMap &materializedLeafMap
6238 )
6239 : OpRewritePattern<constrain::EmitEqualityOp>(ctx), materializedLeaves(materializedLeafMap) {}
6240
6241 LogicalResult
6242 matchAndRewrite(constrain::EmitEqualityOp op, PatternRewriter &rewriter) const override {
6243 PodType podTy = getWholePodEqualityType(op);
6244 if (!podTy) {
6245 return failure();
6246 }
6247
6248 return splitWholePodEmitEquality(op, rewriter, materializedLeaves);
6249 }
6250};
6251
6253static LogicalResult
6254applyGreedily(ModuleOp modOp, RewritePatternSet &&patterns, bool *changed = nullptr) {
6255 return applyPatternsGreedily(
6256 modOp->getRegion(0), std::move(patterns),
6257 GreedyRewriteConfig {.fold = false, .cseConstants = false}, changed
6258 );
6259}
6260
6263static LogicalResult step4(ModuleOp modOp) {
6264 CompatiblePodLeafMaterializationMap materializedLeaves;
6265 RewritePatternSet patterns(modOp.getContext());
6266 patterns.add<
6267 FoldReadAfterWriteInBlockPattern, ReplaceIfReadPattern, LiftPodWritesFromIfBlocksPattern,
6268 LiftPodAccessesFromForLoopPattern, LiftPodAccessesFromWhileLoopPattern,
6269 FoldIfCarriedPodReadAfterWritePattern>(patterns.getContext());
6270 patterns.add<SplitPodInEmitEqualityPattern>(patterns.getContext(), materializedLeaves);
6271
6272 LLVM_DEBUG(llvm::dbgs() << "Begin step 4: refactor pod ops within SCF regions\n";);
6273 return applyGreedily(modOp, std::move(patterns));
6274}
6275
6278static bool applyIfCarriedPodReadAfterWritePatterns(ModuleOp modOp) {
6279 RewritePatternSet patterns(modOp.getContext());
6280 patterns.add<FoldIfCarriedPodReadAfterWritePattern>(patterns.getContext());
6281
6282 bool changed = false;
6283 if (failed(applyGreedily(modOp, std::move(patterns), &changed))) {
6284 return false;
6285 }
6286 return changed;
6287}
6288
6291static size_t podTypeScalarizationWeight(Type type) {
6292 auto podTy = dyn_cast<PodType>(type);
6293 if (!podTy) {
6294 return 0;
6295 }
6296
6297 size_t weight = 1;
6298 for (RecordAttr record : podTy.getRecords()) {
6299 weight += podTypeScalarizationWeight(record.getType());
6300 }
6301 return weight;
6302}
6303
6307static size_t podAllocScalarizationWeight(ModuleOp modOp) {
6308 size_t weight = 0;
6309 modOp.walk([&weight](NewPodOp newPodOp) {
6310 weight += podTypeScalarizationWeight(newPodOp.getType());
6311 });
6312 return weight;
6313}
6314
6316inline static bool isResidualPodLikeType(Type type) {
6317 return splittablePod(type) || splittablePodArray(type);
6318}
6319
6321template <typename TypeRangeLike>
6322static size_t countResidualPodLikeTypes(const TypeRangeLike &types) {
6323 size_t count = 0;
6324 for (Type type : types) {
6325 if (isResidualPodLikeType(type)) {
6326 ++count;
6327 }
6328 }
6329 return count;
6330}
6331
6333static bool isResidualPodPlaceholderCast(UnrealizedConversionCastOp castOp) {
6334 return countResidualPodLikeTypes(castOp.getOperandTypes()) != 0 ||
6335 countResidualPodLikeTypes(castOp.getResultTypes()) != 0;
6336}
6337
6339static size_t countResidualPodIR(ModuleOp modOp) {
6340 size_t count = 0;
6341 modOp.walk([&count](Operation *op) {
6342 if (isa<NewPodOp, ReadPodOp, WritePodOp>(op)) {
6343 ++count;
6344 } else if (auto castOp = dyn_cast<UnrealizedConversionCastOp>(op);
6345 castOp && isResidualPodPlaceholderCast(castOp)) {
6346 ++count;
6347 }
6348
6349 count += countResidualPodLikeTypes(op->getOperandTypes());
6350 count += countResidualPodLikeTypes(op->getResultTypes());
6351
6352 for (Region &region : op->getRegions()) {
6353 for (Block &block : region) {
6354 count += countResidualPodLikeTypes(block.getArgumentTypes());
6355 }
6356 }
6357
6358 if (auto funcDef = dyn_cast<FuncDefOp>(op)) {
6359 FunctionType funcTy = funcDef.getFunctionType();
6360 count += countResidualPodLikeTypes(funcTy.getInputs());
6361 count += countResidualPodLikeTypes(funcTy.getResults());
6362 } else if (auto memberDef = dyn_cast<MemberDefOp>(op)) {
6363 count += isResidualPodLikeType(memberDef.getType()) ? 1 : 0;
6364 }
6365 });
6366 return count;
6367}
6368
6370static LogicalResult rejectRaggedNestedLeafArrayLengths(ModuleOp modOp) {
6371 auto result = modOp.walk([](ArrayLengthOp op) -> WalkResult {
6372 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getArrRef());
6373 if (raggedKind.empty()) {
6374 return WalkResult::advance();
6375 }
6376 return op.emitOpError() << "cannot lower nested " << raggedKind
6377 << " array leaf length after reading an array-of-POD element "
6378 "without per-element shape witnesses";
6379 });
6380 return failure(result.wasInterrupted());
6381}
6382
6384static LogicalResult rejectRaggedNestedLeafArrayEqualities(ModuleOp modOp) {
6385 auto result = modOp.walk([](constrain::EmitEqualityOp op) -> WalkResult {
6386 StringRef raggedKind = getTaggedRaggedNestedLeafKind(op.getLhs());
6387 if (raggedKind.empty()) {
6388 raggedKind = getTaggedRaggedNestedLeafKind(op.getRhs());
6389 }
6390 if (raggedKind.empty()) {
6391 return WalkResult::advance();
6392 }
6393 return op.emitOpError() << "cannot lower nested " << raggedKind
6394 << " array leaf equality after reading an array-of-POD element "
6395 "without per-element shape witnesses";
6396 });
6397 return failure(result.wasInterrupted());
6398}
6399
6401static LogicalResult rejectRaggedNestedLeafArrayContainments(ModuleOp modOp) {
6402 auto result = modOp.walk([](constrain::EmitContainmentOp op) -> WalkResult {
6403 return rejectRaggedNestedLeafContainment(op);
6404 });
6405 return failure(result.wasInterrupted());
6406}
6407
6409static LogicalResult rejectRaggedNestedLeafBoundaryCrossings(ModuleOp modOp) {
6410 auto result = modOp.walk([](Operation *op) -> WalkResult {
6411 return TypeSwitch<Operation *, LogicalResult>(op)
6412 .Case<CallOp>([](CallOp callOp) {
6413 return rejectRaggedNestedLeafBoundaryCrossing(
6414 callOp.getOperation(), callOp.getArgOperands(), "across a function call"
6415 );
6416 })
6417 .Case<ReturnOp>([](ReturnOp returnOp) {
6418 return rejectRaggedNestedLeafBoundaryCrossing(
6419 returnOp.getOperation(), returnOp.getOperands(), "across a function return"
6420 );
6421 })
6422 .Case<scf::ForOp>([](scf::ForOp forOp) {
6423 return rejectRaggedNestedLeafBoundaryCrossing(
6424 forOp.getOperation(), forOp.getInitArgs(), "into an scf.for region"
6425 );
6426 })
6427 .Case<scf::WhileOp>([](scf::WhileOp whileOp) {
6428 return rejectRaggedNestedLeafBoundaryCrossing(
6429 whileOp.getOperation(), whileOp.getInits(), "into an scf.while region"
6430 );
6431 })
6432 .Case<scf::ConditionOp>([](scf::ConditionOp conditionOp) {
6433 return rejectRaggedNestedLeafBoundaryCrossing(
6434 conditionOp.getOperation(), conditionOp.getArgs(), "across an scf.condition boundary"
6435 );
6436 })
6437 .Case<scf::YieldOp>([](scf::YieldOp yieldOp) {
6438 return rejectRaggedNestedLeafBoundaryCrossing(
6439 yieldOp.getOperation(), yieldOp.getOperands(), "across an scf.yield boundary"
6440 );
6441 }).Default([](Operation *) { return success(); });
6442 });
6443 return failure(result.wasInterrupted());
6444}
6445
6447static LogicalResult rejectRemainingRaggedNestedLeafUses(ModuleOp modOp) {
6448 if (failed(rejectRaggedNestedLeafArrayLengths(modOp))) {
6449 return failure();
6450 }
6451 if (failed(rejectRaggedNestedLeafArrayEqualities(modOp))) {
6452 return failure();
6453 }
6454 if (failed(rejectRaggedNestedLeafArrayContainments(modOp))) {
6455 return failure();
6456 }
6457 return rejectRaggedNestedLeafBoundaryCrossings(modOp);
6458}
6459
6460// In MLIR 20.1.8, the `RemoveDeadValues` pass did not handle branching regions as aggressively
6461// as the `RunLivenessAnalysis` that it depends on. The liveness analysis can conclude an `scf.if`
6462// branch is dead and thus its `scf.yield` operand is also dead. However, RDV does not end up
6463// removing that branch but does remove the definition of the `scf.yield` operand, leaving a NULL
6464// operand in the `scf.yield` op.
6465//
6466// The bug is fixed in MLIR 22.1.0 but there is a temporary solution: before running the RDV pass,
6467// run "Sparse Conditional Constant Propagation" (SCCP) plus a custom pass that folds `scf.if` with
6468// static conditions.
6469//
6470// Portions adapted from mlir/lib/Dialect/SCF/IR/SCF.cpp
6471// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
6472// See https://llvm.org/LICENSE.txt for license information.
6473// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6475
6476static void replaceOpWithRegion(
6477 PatternRewriter &rewriter, Operation *op, Region &region, ValueRange blockArgs = {}
6478) {
6479 assert(llvm::hasSingleElement(region) && "expected single-region block");
6480 Block *block = &region.front();
6481 Operation *terminator = block->getTerminator();
6482 ValueRange results = terminator->getOperands();
6483 rewriter.inlineBlockBefore(block, op, blockArgs);
6484 rewriter.replaceOp(op, results);
6485 rewriter.eraseOp(terminator);
6486}
6487
6488struct RemoveStaticCondition : public OpRewritePattern<scf::IfOp> {
6489 using OpRewritePattern<scf::IfOp>::OpRewritePattern;
6490
6491 LogicalResult matchAndRewrite(scf::IfOp op, PatternRewriter &rewriter) const override {
6492 BoolAttr condition;
6493 if (!matchPattern(op.getCondition(), m_Constant(&condition))) {
6494 return failure();
6495 }
6496 if (condition.getValue()) {
6497 replaceOpWithRegion(rewriter, op, op.getThenRegion());
6498 } else if (!op.getElseRegion().empty()) {
6499 replaceOpWithRegion(rewriter, op, op.getElseRegion());
6500 } else {
6501 rewriter.eraseOp(op);
6502 }
6503 return success();
6504 }
6505};
6506
6507struct FlattenStaticIfPass : PassWrapper<FlattenStaticIfPass, OperationPass<ModuleOp>> {
6508 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(FlattenStaticIfPass)
6509
6510 void runOnOperation() override {
6511 RewritePatternSet patterns(&getContext());
6512 patterns.add<RemoveStaticCondition>(&getContext());
6513
6514 if (failed(applyPatternsGreedily(
6515 getOperation(), std::move(patterns),
6516 GreedyRewriteConfig {.fold = false, .cseConstants = false}
6517 ))) {
6518 signalPassFailure();
6519 }
6520 }
6521};
6522
6523void add(OpPassManager &pm) {
6524 pm.addPass(createSCCPPass());
6525 pm.addPass(std::make_unique<FlattenStaticIfPass>());
6526}
6527
6528} // namespace temp_fix_pre_mlir_22
6529
6531class PassImpl : public llzk::pod::impl::PodToScalarPassBase<PassImpl> {
6532 using Base = PodToScalarPassBase<PassImpl>;
6533 using Base::Base;
6534
6535 LogicalResult runScalarizeAndCleanupPipeline(ModuleOp module) {
6536 // 1. Use SROA (Destructurable* interfaces) to split each pod with `N` records into `N` pods
6537 // with 1 record each. This is necessary because the mem2reg pass cannot deal with splitting
6538 // up memory, i.e., it can only convert scalar memory access into SSA values.
6539 // 2. The mem2reg pass converts the size 1 pod allocations and accesses into SSA values.
6540 OpPassManager scalarizePM(ModuleOp::getOperationName());
6541 scalarizePM.addPass(createSpecializedSROAPass<NewPodOp>());
6542 scalarizePM.addPass(createSpecializedMem2RegPass<NewPodOp>());
6543
6544 // Cleanup allocations made dead by memory promotion and other dead SSA values.
6545 OpPassManager cleanupPM(ModuleOp::getOperationName());
6548 .allocatorOpName = CreateArrayOp::getOperationName().str()
6549 }
6550 ));
6552 RemoveUnusedDiscardableAllocationsPassOptions {
6553 .allocatorOpName = NewPodOp::getOperationName().str()
6554 }
6555 ));
6556 temp_fix_pre_mlir_22::add(cleanupPM);
6557 cleanupPM.addPass(createRemoveDeadValuesWorkaroundPass());
6558
6559 size_t podAllocWeight = podAllocScalarizationWeight(module);
6560 while (podAllocWeight != 0) {
6561 if (failed(runPipeline(scalarizePM, module))) {
6562 return failure();
6563 }
6564
6565 // SROA+mem2reg can expose `scf.if`-carried POD values that become redundant after a
6566 // same-record write from another `scf.if` result. Fold those reads and clean up before
6567 // checking convergence.
6568 bool foldedIfCarriedRead = applyIfCarriedPodReadAfterWritePatterns(module);
6569 if (failed(runPipeline(cleanupPM, module))) {
6570 return failure();
6571 }
6572
6573 // Nested PODs can become visible only after an outer single-record POD has been promoted,
6574 // and SROA can transiently increase allocation count while splitting aggregates. Keep
6575 // iterating until the allocation-weight heuristic reaches a fixed point.
6576 size_t nextPodAllocWeight = podAllocScalarizationWeight(module);
6577 if (!foldedIfCarriedRead && nextPodAllocWeight == podAllocWeight) {
6578 break;
6579 }
6580 podAllocWeight = nextPodAllocWeight;
6581 }
6582
6583 return success();
6584 }
6585
6586 void runOnOperation() override {
6587 ModuleOp module = getOperation();
6588 if (failed(step0(module))) {
6589 return signalPassFailure();
6590 }
6591 LLVM_DEBUG({
6592 llvm::dbgs() << "After step 0:\n";
6593 module.dump();
6594 });
6595
6596 size_t previousResidualCount = std::numeric_limits<size_t>::max();
6597 while (true) {
6598 // This is divided into 2 steps to simplify the implementation for member-related ops. The
6599 // issue is that the conversions for member read/write expect the mapping of record name to
6600 // member name+type to already be populated for the referenced member (although this could be
6601 // computed on demand if desired but it complicates the implementation a bit).
6602 SymbolTableCollection symTables;
6603 MemberReplacementMap memberRepMap;
6604 if (failed(step1(module, symTables, memberRepMap))) {
6605 return signalPassFailure();
6606 }
6607 LLVM_DEBUG({
6608 llvm::dbgs() << "After step 1:\n";
6609 module.dump();
6610 });
6611
6612 if (failed(step2(module, symTables, memberRepMap))) {
6613 return signalPassFailure();
6614 }
6615 LLVM_DEBUG({
6616 llvm::dbgs() << "After step 2:\n";
6617 module.dump();
6618 });
6619
6620 if (failed(step3(module, symTables, memberRepMap))) {
6621 return signalPassFailure();
6622 }
6623 LLVM_DEBUG({
6624 llvm::dbgs() << "After step 3:\n";
6625 module.dump();
6626 });
6627
6628 if (failed(rejectRemainingRaggedNestedLeafUses(module))) {
6629 signalPassFailure();
6630 return;
6631 }
6632
6633 if (failed(step4(module))) {
6634 return signalPassFailure();
6635 }
6636 LLVM_DEBUG({
6637 llvm::dbgs() << "After step 4:\n";
6638 module.dump();
6639 });
6640
6641 if (failed(runScalarizeAndCleanupPipeline(module))) {
6642 signalPassFailure();
6643 return;
6644 }
6645 LLVM_DEBUG({
6646 llvm::dbgs() << "After SROA+Mem2Reg pipeline:\n";
6647 module.dump();
6648 });
6649
6650 if (failed(rejectRemainingRaggedNestedLeafUses(module))) {
6651 signalPassFailure();
6652 return;
6653 }
6654
6655 size_t residualCount = countResidualPodIR(module);
6656 if (residualCount == 0) {
6657 break;
6658 }
6659 if (residualCount >= previousResidualCount) {
6660 std::string residualIR;
6661 llvm::raw_string_ostream os(residualIR);
6662 module.print(os);
6663 module.emitError() << "llzk-pod-to-scalar left residual pod IR after reaching a fixpoint ("
6664 << residualCount << " residual pod-like items remain)\n"
6665 << residualIR;
6666 signalPassFailure();
6667 return;
6668 }
6669 previousResidualCount = residualCount;
6670 }
6671 }
6672};
6673
6674} // 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
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:564
::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::FunctionType getFunctionType()
Definition Ops.cpp.inc:984
void setArgAttrsAttr(::mlir::ArrayAttr attr)
Definition Ops.h.inc:746
::llvm::ArrayRef<::mlir::Type > getArgumentTypes()
Required by FunctionOpInterface.
Definition Ops.h.inc:883
void setResAttrsAttr(::mlir::ArrayAttr attr)
Definition Ops.h.inc:750
::mlir::ArrayAttr getArgAttrsAttr()
Definition Ops.h.inc:726
void setFunctionType(::mlir::FunctionType attrValue)
Definition Ops.cpp.inc:1003
::llvm::ArrayRef<::mlir::Type > getResultTypes()
Required by FunctionOpInterface.
Definition Ops.h.inc:887
::mlir::Region & getBody()
Definition Ops.h.inc:703
::mlir::ArrayAttr getResAttrsAttr()
Definition Ops.h.inc:731
::mlir::MutableOperandRange getOperandsMutable()
Definition Ops.cpp.inc:1169
::mlir::Operation::operand_range getOperands()
Definition Ops.h.inc:1024
::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:1329
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()
ExpressionValue add(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
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:94
ArrayType flattenArrayElementType(ArrayType outerArrTy, Type elementType)
std::unique_ptr< SpecializedMem2Reg< AllocOpTy > > createSpecializedMem2RegPass()
ConstantCapture m_Constant()
Definition Matchers.h:89
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:64
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