From 5ce0a5984b880705c838577d1c88e713ea89bf32 Mon Sep 17 00:00:00 2001 From: Joao Roberto Date: Tue, 11 Aug 2026 13:17:11 -0300 Subject: [PATCH] Add compact strided TypeTree ranges for large arrays Dense per-element offsets waste memory and are dropped above MaxTypeOffset (default 500), so mixed aggregates like `{ i64, [f32; 1000] }` could not describe the final array element. Introduce `start+stride*count` index syntax and first-class StridedTypeRange storage so frontends can emit compact homogeneous array regions that remain queryable past the offset cap. --- enzyme/Enzyme/CApi.cpp | 12 + enzyme/Enzyme/CApi.h | 6 + enzyme/Enzyme/TypeAnalysis/TypeTree.h | 375 ++++++++++++++++++++++- enzyme/test/TypeAnalysis/stridedrange.ll | 24 ++ 4 files changed, 404 insertions(+), 13 deletions(-) create mode 100644 enzyme/test/TypeAnalysis/stridedrange.ll diff --git a/enzyme/Enzyme/CApi.cpp b/enzyme/Enzyme/CApi.cpp index 9b584f17a633..a460ad0f3ea1 100644 --- a/enzyme/Enzyme/CApi.cpp +++ b/enzyme/Enzyme/CApi.cpp @@ -933,6 +933,18 @@ void EnzymeTypeTreeInsertEq(CTypeTreeRef CTT, const int64_t *indices, } ((TypeTree *)CTT)->insert(seq, eunwrap(ct, *unwrap(ctx))); } +void EnzymeTypeTreeInsertRangeEq(CTypeTreeRef CTT, const int64_t *indices, + size_t len, size_t rangePos, int64_t stride, + int64_t count, CConcreteType ct, + LLVMContextRef ctx) { + std::vector seq; + for (size_t i = 0; i < len; i++) { + seq.push_back(indices[i]); + } + ((TypeTree *)CTT) + ->insertRange(seq, rangePos, (int)stride, (int)count, + eunwrap(ct, *unwrap(ctx))); +} const char *EnzymeTypeTreeToString(CTypeTreeRef src) { std::string tmp = ((TypeTree *)src)->str(); char *cstr = new char[tmp.length() + 1]; diff --git a/enzyme/Enzyme/CApi.h b/enzyme/Enzyme/CApi.h index 9d70a02b7cf0..8dfd33020c35 100644 --- a/enzyme/Enzyme/CApi.h +++ b/enzyme/Enzyme/CApi.h @@ -98,6 +98,12 @@ void EnzymeTypeTreeShiftIndiciesEq(CTypeTreeRef dst, const char *datalayout, uint64_t addOffset); void EnzymeTypeTreeInsertEq(CTypeTreeRef dst, const int64_t *indices, size_t len, CConcreteType ct, LLVMContextRef ctx); +/// Insert a compact strided range. `indices[rangePos]` is the start offset; +/// the range covers start + k*stride for k in 0..count. +void EnzymeTypeTreeInsertRangeEq(CTypeTreeRef dst, const int64_t *indices, + size_t len, size_t rangePos, int64_t stride, + int64_t count, CConcreteType ct, + LLVMContextRef ctx); const char *EnzymeTypeTreeToString(CTypeTreeRef src); void EnzymeTypeTreeToStringFree(const char *cstr); diff --git a/enzyme/Enzyme/TypeAnalysis/TypeTree.h b/enzyme/Enzyme/TypeAnalysis/TypeTree.h index 10ee9cd5a539..a539f30469f5 100644 --- a/enzyme/Enzyme/TypeAnalysis/TypeTree.h +++ b/enzyme/Enzyme/TypeAnalysis/TypeTree.h @@ -33,6 +33,7 @@ #include "llvm/Support/ErrorHandling.h" #include "llvm/Support/raw_ostream.h" +#include #include #include #include @@ -61,6 +62,69 @@ static inline std::string to_string(const std::vector x) { return out; } +/// Compact strided byte-offset range for homogeneous array regions. +/// +/// Wire form uses `start+stride*count` in place of a single index component, +/// e.g. `{[-1,8+4*1000]:Float@float}` means floats at bytes +/// `8 + 4*k` for `k in 0..1000` after one pointer dereference. This avoids +/// materializing `count` map keys and survives `MaxTypeOffset`, which would +/// otherwise drop dense offsets above the analysis cap. +struct StridedTypeRange { + std::vector Indices; // Indices[RangePos] holds the start offset + size_t RangePos = 0; + int Stride = 0; + int Count = 0; + ConcreteType CT = BaseType::Unknown; + + int start() const { return Indices[RangePos]; } + int last() const { return start() + Stride * (Count - 1); } + + bool offsetInRange(int idx) const { + if (idx < start() || Stride <= 0 || Count <= 0) + return false; + int delta = idx - start(); + if (delta % Stride != 0) + return false; + int k = delta / Stride; + return k >= 0 && k < Count; + } + + /// True if `Seq` equals `Indices` on non-range positions and lands inside + /// the strided span at `RangePos`. + bool matches(const std::vector &Seq) const { + if (Seq.size() != Indices.size()) + return false; + for (size_t i = 0, e = Seq.size(); i < e; ++i) { + if (i == RangePos) { + if (!offsetInRange(Seq[i])) + return false; + } else if (Indices[i] != Seq[i]) { + return false; + } + } + return true; + } + + std::string indexString() const { + std::string out = "["; + for (unsigned i = 0; i < Indices.size(); ++i) { + if (i != 0) + out += ","; + if (i == RangePos) { + out += std::to_string(Indices[i]); + out += "+"; + out += std::to_string(Stride); + out += "*"; + out += std::to_string(Count); + } else { + out += std::to_string(Indices[i]); + } + } + out += "]"; + return out; + } +}; + class TypeTree; typedef std::shared_ptr TypeResult; @@ -74,6 +138,8 @@ class TypeTree : public std::enable_shared_from_this { // mapping of known indices to type if one exists ConcreteTypeMapType mapping; std::vector minIndices; + // Compact homogeneous array regions (see StridedTypeRange). + std::vector ranges; public: TypeTree() {} @@ -99,6 +165,9 @@ class TypeTree : public std::enable_shared_from_this { str = str.substr(1); std::vector idxs; + size_t rangePos = (size_t)-1; + int rangeStride = 0; + int rangeCount = 0; while (true) { while (str[0] == ' ') str = str.substr(1); @@ -111,6 +180,25 @@ class TypeTree : public std::enable_shared_from_this { bool failed = str.consumeInteger(10, idx); (void)failed; assert(!failed); + + // Optional compact range: start+stride*count + if (!str.empty() && str[0] == '+') { + str = str.substr(1); + int stride = 0; + int count = 0; + failed = str.consumeInteger(10, stride); + assert(!failed); + assert(!str.empty() && str[0] == '*'); + str = str.substr(1); + failed = str.consumeInteger(10, count); + assert(!failed); + assert(rangePos == (size_t)-1 && "at most one ranged index"); + assert(stride > 0 && count > 0); + rangePos = idxs.size(); + rangeStride = stride; + rangeCount = count; + } + idxs.push_back(idx); while (str[0] == ' ') @@ -146,16 +234,38 @@ class TypeTree : public std::enable_shared_from_this { str = str.substr(endval); ConcreteType CT(tystr, ctx); - Result.mapping.emplace(idxs, CT); - if (Result.minIndices.size() < idxs.size()) { - for (size_t i = Result.minIndices.size(), end = idxs.size(); i < end; - ++i) { - Result.minIndices.push_back(idxs[i]); + if (rangePos != (size_t)-1) { + StridedTypeRange R; + R.Indices = idxs; + R.RangePos = rangePos; + R.Stride = rangeStride; + R.Count = rangeCount; + R.CT = CT; + Result.ranges.push_back(R); + int start = idxs[rangePos]; + if (Result.minIndices.size() < idxs.size()) { + for (size_t i = Result.minIndices.size(), end = idxs.size(); i < end; + ++i) { + Result.minIndices.push_back(i == rangePos ? start : idxs[i]); + } + } + for (size_t i = 0, end = idxs.size(); i < end; ++i) { + int v = (i == rangePos) ? start : idxs[i]; + if (v < Result.minIndices[i]) + Result.minIndices[i] = v; + } + } else { + Result.mapping.emplace(idxs, CT); + if (Result.minIndices.size() < idxs.size()) { + for (size_t i = Result.minIndices.size(), end = idxs.size(); i < end; + ++i) { + Result.minIndices.push_back(idxs[i]); + } + } + for (size_t i = 0, end = idxs.size(); i < end; ++i) { + if (idxs[i] < Result.minIndices[i]) + Result.minIndices[i] = idxs[i]; } - } - for (size_t i = 0, end = idxs.size(); i < end; ++i) { - if (idxs[i] < Result.minIndices[i]) - Result.minIndices[i] = idxs[i]; } while (str[0] == ' ') @@ -178,6 +288,10 @@ class TypeTree : public std::enable_shared_from_this { auto Found0 = mapping.find(Seq); if (Found0 != mapping.end()) return Found0->second; + for (const auto &R : ranges) { + if (R.matches(Seq)) + return R.CT; + } size_t Len = Seq.size(); if (Len == 0) return BaseType::Unknown; @@ -213,6 +327,26 @@ class TypeTree : public std::enable_shared_from_this { return Found->second; } } + // Check ranges against -1-generalized sequences as well: a query for a + // concrete offset should hit `[-1, start+stride*count]` ranges. + for (const auto &R : ranges) { + if (R.Indices.size() != Seq.size()) + continue; + bool ok = true; + for (size_t j = 0; j < Seq.size(); ++j) { + if (j == R.RangePos) { + if (!R.offsetInRange(Seq[j])) { + ok = false; + break; + } + } else if (R.Indices[j] != Seq[j] && R.Indices[j] != -1) { + ok = false; + break; + } + } + if (ok) + return R.CT; + } return BaseType::Unknown; } @@ -426,8 +560,72 @@ class TypeTree : public std::enable_shared_from_this { return true; } + /// Insert a compact strided range. `Seq[RangePos]` is the start offset; + /// the range covers `start + k*Stride` for `k in 0..Count`. + /// Unlike `insert`, ranges are not subject to `MaxTypeOffset` pruning. + bool insertRange(const std::vector Seq, size_t RangePos, int Stride, + int Count, ConcreteType CT) { + assert(RangePos < Seq.size()); + assert(Stride > 0); + assert(Count > 0); + if (CT == ConcreteType(BaseType::Unknown)) + return false; + if (Count == 1) { + std::vector single = Seq; + return insert(single, CT); + } + for (const auto &existing : ranges) { + if (existing.Indices == Seq && existing.RangePos == RangePos && + existing.Stride == Stride && existing.Count == Count) { + if (existing.CT == CT) + return false; + llvm::errs() << "inserting conflicting range into : " << str() + << " with " << to_string(Seq) << " stride=" << Stride + << " count=" << Count << " of " << CT.str() << "\n"; + llvm_unreachable("illegal range insertion"); + } + } + StridedTypeRange R; + R.Indices = Seq; + R.RangePos = RangePos; + R.Stride = Stride; + R.Count = Count; + R.CT = CT; + ranges.push_back(R); + + int start = Seq[RangePos]; + if (minIndices.size() < Seq.size()) { + for (size_t i = minIndices.size(), end = Seq.size(); i < end; ++i) + minIndices.push_back(i == RangePos ? start : Seq[i]); + } + for (size_t i = 0, end = Seq.size(); i < end; ++i) { + int v = (i == RangePos) ? start : Seq[i]; + if (v < minIndices[i]) + minIndices[i] = v; + } + return true; + } + /// How this TypeTree compares with another - bool operator<(const TypeTree &vd) const { return mapping < vd.mapping; } + bool operator<(const TypeTree &vd) const { + if (mapping != vd.mapping) + return mapping < vd.mapping; + if (ranges.size() != vd.ranges.size()) + return ranges.size() < vd.ranges.size(); + for (size_t i = 0; i < ranges.size(); ++i) { + if (ranges[i].Indices != vd.ranges[i].Indices) + return ranges[i].Indices < vd.ranges[i].Indices; + if (ranges[i].RangePos != vd.ranges[i].RangePos) + return ranges[i].RangePos < vd.ranges[i].RangePos; + if (ranges[i].Stride != vd.ranges[i].Stride) + return ranges[i].Stride < vd.ranges[i].Stride; + if (ranges[i].Count != vd.ranges[i].Count) + return ranges[i].Count < vd.ranges[i].Count; + if (ranges[i].CT != vd.ranges[i].CT) + return ranges[i].CT < vd.ranges[i].CT; + } + return false; + } /// Whether this TypeTree contains any information bool isKnown() const { @@ -438,7 +636,7 @@ class TypeTree : public std::enable_shared_from_this { assert(pair.second.isKnown()); } #endif - return mapping.size() != 0; + return mapping.size() != 0 || ranges.size() != 0; } /// Whether this TypeTree knows any non-pointer information @@ -454,7 +652,7 @@ class TypeTree : public std::enable_shared_from_this { } return true; } - return false; + return !ranges.empty(); } /// Select only the Integer ConcreteTypes @@ -508,6 +706,14 @@ class TypeTree : public std::enable_shared_from_this { Result.mapping.insert( std::pair, ConcreteType>(Vec, pair.second)); } + for (const auto &R : ranges) { + if (R.Indices.size() == EnzymeMaxTypeDepth) + continue; + StridedTypeRange Next = R; + Next.Indices.insert(Next.Indices.begin(), Off); + Next.RangePos = R.RangePos + 1; + Result.ranges.push_back(Next); + } return Result; } @@ -542,6 +748,33 @@ class TypeTree : public std::enable_shared_from_this { Result.orIn(next, pair.second); } } + for (const auto &R : ranges) { + assert(!R.Indices.empty()); + if (R.RangePos == 0) { + // Outermost index is the ranged dimension: keep the range only when it + // covers offset 0 (or would via -1, which ranges do not use). + if (!R.offsetInRange(0)) + continue; + StridedTypeRange Next = R; + Next.Indices.erase(Next.Indices.begin()); + if (Next.RangePos > 0) + Next.RangePos--; + // After peeling a ranged outermost index that matched 0, the remaining + // path is a concrete (non-range) lookup at depth 0 of the child. + if (Next.Indices.empty()) { + Result.orIn(std::vector{}, Next.CT); + } else if (Next.RangePos >= Next.Indices.size()) { + Result.orIn(Next.Indices, Next.CT); + } else { + Result.ranges.push_back(Next); + } + } else if (R.Indices[0] == -1 || R.Indices[0] == 0) { + StridedTypeRange Next = R; + Next.Indices.erase(Next.Indices.begin()); + Next.RangePos = R.RangePos - 1; + Result.ranges.push_back(Next); + } + } return Result; } @@ -1103,6 +1336,67 @@ class TypeTree : public std::enable_shared_from_this { // Resize minIndices down if we dropped any higher-depth indices for being // out of scope. Result.minIndices.resize(maxInsertedDepth); + + // Shift compact strided ranges. Ranges stay compact (no per-element + // materialization) and are not discarded by MaxTypeOffset. + for (const auto &R : ranges) { + if (R.Indices.empty()) + continue; + if (R.RangePos == 0) { + int nextStart = R.start(); + if (nextStart < offset) + continue; + nextStart -= offset; + if (maxSize != -1 && nextStart >= maxSize) + continue; + // Keep the range only if at least one element remains in-window. + int remaining = R.Count; + if (maxSize != -1) { + // Largest k with start + k*stride < maxSize + remaining = (maxSize - nextStart + R.Stride - 1) / R.Stride; + if (remaining > R.Count) + remaining = R.Count; + if (remaining <= 0) + continue; + } + nextStart += addOffset; + StridedTypeRange Next = R; + Next.Indices[0] = nextStart; + Next.Count = remaining; + Result.ranges.push_back(Next); + if (nextStart < Result.minIndices[0]) + Result.minIndices[0] = nextStart; + if (Next.Indices.size() > maxInsertedDepth) + maxInsertedDepth = Next.Indices.size(); + } else { + // Non-outermost ranged dimension: keep if the outermost index shifts + // into range, mirroring concrete-map handling. + int next0 = R.Indices[0]; + if (next0 != -1) { + if (next0 < offset) + continue; + next0 -= offset; + if (maxSize != -1 && next0 >= maxSize) + continue; + next0 += addOffset; + } else if (addOffset != 0 && maxSize == -1) { + next0 = addOffset; + } else if (maxSize != -1) { + // -1 outermost with finite maxSize would need expansion; keep as-is + // with addOffset applied when possible. + if (addOffset != 0) + next0 = addOffset; + } + StridedTypeRange Next = R; + Next.Indices[0] = next0; + Result.ranges.push_back(Next); + if (next0 != -1 && next0 < Result.minIndices[0]) + Result.minIndices[0] = next0; + if (Next.Indices.size() > maxInsertedDepth) + maxInsertedDepth = Next.Indices.size(); + } + } + Result.minIndices.resize(maxInsertedDepth); return Result; } @@ -1160,7 +1454,19 @@ class TypeTree : public std::enable_shared_from_this { } /// Chceck equality of two TypeTrees - bool operator==(const TypeTree &RHS) const { return mapping == RHS.mapping; } + bool operator==(const TypeTree &RHS) const { + if (mapping != RHS.mapping || ranges.size() != RHS.ranges.size()) + return false; + for (size_t i = 0; i < ranges.size(); ++i) { + if (ranges[i].Indices != RHS.ranges[i].Indices || + ranges[i].RangePos != RHS.ranges[i].RangePos || + ranges[i].Stride != RHS.ranges[i].Stride || + ranges[i].Count != RHS.ranges[i].Count || + ranges[i].CT != RHS.ranges[i].CT) + return false; + } + return true; + } /// Set this to another TypeTree, returning if this was changed bool operator=(const TypeTree &RHS) { @@ -1316,6 +1622,26 @@ class TypeTree : public std::enable_shared_from_this { for (auto &pair : RHS.mapping) { changed |= checkedOrIn(pair.first, pair.second, PointerIntSame, LegalOr); } + for (const auto &R : RHS.ranges) { + bool found = false; + for (auto &existing : ranges) { + if (existing.Indices == R.Indices && existing.RangePos == R.RangePos && + existing.Stride == R.Stride && existing.Count == R.Count) { + found = true; + if (existing.CT == R.CT) + break; + bool sub = existing.CT.checkedOrIn(R.CT, PointerIntSame, LegalOr); + if (!LegalOr) + return changed; + changed |= sub; + break; + } + } + if (!found) { + ranges.push_back(R); + changed = true; + } + } return changed; } @@ -1465,6 +1791,29 @@ class TypeTree : public std::enable_shared_from_this { out += "]:" + pair.second.str(); first = false; } + // Emit ranges after concrete mappings, sorted for stability. + std::vector ordered; + ordered.reserve(ranges.size()); + for (const auto &R : ranges) + ordered.push_back(&R); + std::sort(ordered.begin(), ordered.end(), + [](const StridedTypeRange *A, const StridedTypeRange *B) { + if (A->Indices != B->Indices) + return A->Indices < B->Indices; + if (A->RangePos != B->RangePos) + return A->RangePos < B->RangePos; + if (A->Stride != B->Stride) + return A->Stride < B->Stride; + return A->Count < B->Count; + }); + for (const auto *R : ordered) { + if (!first) { + out += ", "; + } + out += R->indexString(); + out += ":" + R->CT.str(); + first = false; + } out += "}"; return out; } diff --git a/enzyme/test/TypeAnalysis/stridedrange.ll b/enzyme/test/TypeAnalysis/stridedrange.ll new file mode 100644 index 000000000000..5b8b9a76e4de --- /dev/null +++ b/enzyme/test/TypeAnalysis/stridedrange.ll @@ -0,0 +1,24 @@ +; RUN: if [ %llvmver -lt 16 ]; then %opt < %s %loadEnzyme -print-type-analysis -type-analysis-func=caller -o /dev/null | FileCheck %s; fi +; RUN: %opt < %s %newLoadEnzyme -passes="print-type-analysis" -type-analysis-func=caller -S | FileCheck %s + +; Verify compact strided TypeTree ranges parse and remain queryable past +; MaxTypeOffset (default 500). The attribute encodes floats at +; 8 + 4*k for k in 0..1000 (last element byte 4004) beside an i64 header. + +target datalayout = "e-m:e-i64:64-f80:128-n8:16:32:64-S128" +target triple = "x86_64-unknown-linux-gnu" + +define float @caller(ptr "enzyme_type"="{[-1]:Pointer, [-1,0]:Integer, [-1,8+4*1000]:Float@float}" %c) { +entry: + %h = load i64, ptr %c, align 8 + %data = getelementptr inbounds i8, ptr %c, i64 8 + %first = load float, ptr %data, align 4 + %lastp = getelementptr inbounds i8, ptr %c, i64 4004 + %last = load float, ptr %lastp, align 4 + %sum = fadd float %first, %last + ret float %sum +} + +; CHECK: caller - {} |{[-1]:Pointer, [-1,0]:Integer, [-1,8+4*1000]:Float@float}:{} +; CHECK: ptr %c: {[-1]:Pointer, [-1,0]:Integer, [-1,8+4*1000]:Float@float} +; CHECK: %last = load float, ptr %lastp, align 4: {[-1]:Float@float}