[clang] [NFC][clang][bytecode][HLSL] Refactor HLSL helper functions to use a common visitor class (PR #194393)

via cfe-commits cfe-commits at lists.llvm.org
Mon Apr 27 07:57:17 PDT 2026


llvmbot wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-clang

Author: Deric C. (Icohedron)

<details>
<summary>Changes</summary>

Addressing the code repetition concerns pointed out by @<!-- -->shafik in https://github.com/llvm/llvm-project/pull/189126#discussion_r3141480667

This PR refactors `emitHLSLAggregateSplat`, `countHLSLFlatElements`, `emitHLSLFlattenAggregate`, and `emitHLSLConstructAggregate` to reduce code repetition by creating a common `HLSLElementStoreVisitor` class for visiting HLSL aggregate types.

Assisted-by: Claude Opus 4.6

---

Patch is 28.32 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/194393.diff


2 Files Affected:

- (modified) clang/lib/AST/ByteCode/Compiler.cpp (+231-365) 
- (modified) clang/lib/AST/ByteCode/Compiler.h (+106-7) 


``````````diff
diff --git a/clang/lib/AST/ByteCode/Compiler.cpp b/clang/lib/AST/ByteCode/Compiler.cpp
index d1878cbedae58..a1f46f0631fbf 100644
--- a/clang/lib/AST/ByteCode/Compiler.cpp
+++ b/clang/lib/AST/ByteCode/Compiler.cpp
@@ -8018,119 +8018,235 @@ bool Compiler<Emitter>::emitBuiltinBitCast(const CastExpr *E) {
   return true;
 }
 
-/// Replicate a scalar value into every scalar element of an aggregate.
-/// The scalar is stored in a local at \p SrcOffset and a pointer to the
-/// destination must be on top of the interpreter stack. Each element receives
-/// the scalar, cast to its own type.
+namespace clang {
+namespace interp {
+
+/// Visitor that stores values into an HLSL destination type.
+/// ProduceValue must leave exactly one value of the requested type on the
+/// interpreter stack; the visitor then stores it with the appropriate
+/// InitElem / InitField / etc. opcode.
 template <class Emitter>
-bool Compiler<Emitter>::emitHLSLAggregateSplat(PrimType SrcT,
-                                               unsigned SrcOffset,
-                                               QualType DestType,
-                                               const Expr *E) {
-  // Vectors and matrices are treated as flat sequences of elements.
-  unsigned NumElems = 0;
-  QualType ElemType;
-  if (const auto *VT = DestType->getAs<VectorType>()) {
-    NumElems = VT->getNumElements();
-    ElemType = VT->getElementType();
-  } else if (const auto *MT = DestType->getAs<ConstantMatrixType>()) {
-    NumElems = MT->getNumElementsFlattened();
-    ElemType = MT->getElementType();
-  }
-  if (NumElems > 0) {
-    PrimType ElemT = classifyPrim(ElemType);
-    for (unsigned I = 0; I != NumElems; ++I) {
-      if (!this->emitGetLocal(SrcT, SrcOffset, E))
-        return false;
-      if (!this->emitPrimCast(SrcT, ElemT, ElemType, E))
-        return false;
-      if (!this->emitInitElem(ElemT, I, E))
-        return false;
-    }
-    return true;
+class HLSLElementStoreVisitor
+    : public Compiler<Emitter>::template HLSLAggregateVisitor<
+          HLSLElementStoreVisitor<Emitter>> {
+  using VisitorBase = typename Compiler<Emitter>::template HLSLAggregateVisitor<
+      HLSLElementStoreVisitor<Emitter>>;
+
+public:
+  HLSLElementStoreVisitor(
+      Compiler<Emitter> &C,
+      llvm::function_ref<bool(PrimType, QualType, const Expr *)> ProduceValue,
+      const Expr *E)
+      : VisitorBase(C), ProduceValue(ProduceValue), E(E) {}
+
+  bool visitScalarElem(QualType ElemType, PrimType ElemT, unsigned I) {
+    if (!ProduceValue(ElemT, ElemType, E))
+      return false;
+    return this->C.emitInitElem(ElemT, I, E);
   }
 
-  // Arrays: primitive elements are filled directly; composite elements
-  // require recursion into each sub-aggregate.
-  if (const auto *AT = DestType->getAsArrayTypeUnsafe()) {
-    const auto *CAT = cast<ConstantArrayType>(AT);
-    QualType ArrElemType = CAT->getElementType();
-    unsigned ArrSize = CAT->getZExtSize();
+  bool visitArrayComposite(QualType ElemType, unsigned I) {
+    if (!this->C.emitConstUint32(I, E))
+      return false;
+    if (!this->C.emitArrayElemPtrUint32(E))
+      return false;
+    if (!this->visit(ElemType))
+      return false;
+    return this->C.emitFinishInitPop(E);
+  }
 
-    if (OptPrimType ElemT = classify(ArrElemType)) {
-      for (unsigned I = 0; I != ArrSize; ++I) {
-        if (!this->emitGetLocal(SrcT, SrcOffset, E))
-          return false;
-        if (!this->emitPrimCast(SrcT, *ElemT, ArrElemType, E))
-          return false;
-        if (!this->emitInitElem(*ElemT, I, E))
-          return false;
-      }
-    } else {
-      for (unsigned I = 0; I != ArrSize; ++I) {
-        if (!this->emitConstUint32(I, E))
-          return false;
-        if (!this->emitArrayElemPtrUint32(E))
-          return false;
-        if (!emitHLSLAggregateSplat(SrcT, SrcOffset, ArrElemType, E))
-          return false;
-        if (!this->emitFinishInitPop(E))
-          return false;
-      }
-    }
+  bool visitBase(QualType BaseType, const Record::Base *B) {
+    if (!this->C.emitGetPtrBase(B->Offset, E))
+      return false;
+    if (!this->visit(BaseType))
+      return false;
+    return this->C.emitFinishInitPop(E);
+  }
+
+  bool visitField(QualType FieldType, PrimType FieldT, const Record::Field *F) {
+    if (!ProduceValue(FieldT, FieldType, E))
+      return false;
+    if (F->isBitField())
+      return this->C.emitInitBitField(FieldT, F->Offset, F->bitWidth(), E);
+    return this->C.emitInitField(FieldT, F->Offset, E);
+  }
+
+  bool visitFieldComposite(QualType FieldType, const Record::Field *F) {
+    if (!this->C.emitGetPtrField(F->Offset, E))
+      return false;
+    if (!this->visit(FieldType))
+      return false;
+    return this->C.emitPopPtr(E);
+  }
+
+private:
+  /// Non-owning — the visitor must not outlive the callable passed at
+  /// construction.
+  llvm::function_ref<bool(PrimType, QualType, const Expr *)> ProduceValue;
+  const Expr *E;
+};
+
+/// Visitor that counts the total number of scalar elements in an HLSL type.
+template <class Emitter>
+class HLSLFlatElementCounter
+    : public Compiler<Emitter>::template HLSLAggregateVisitor<
+          HLSLFlatElementCounter<Emitter>> {
+  using VisitorBase = typename Compiler<Emitter>::template HLSLAggregateVisitor<
+      HLSLFlatElementCounter<Emitter>>;
+
+public:
+  explicit HLSLFlatElementCounter(Compiler<Emitter> &C) : VisitorBase(C) {}
+
+  unsigned getCount() const { return Count; }
+
+  bool visitScalarElem(QualType, PrimType, unsigned) {
+    ++Count;
+    return true;
+  }
+  bool visitArrayComposite(QualType ElemType, unsigned) {
+    return this->visit(ElemType);
+  }
+  bool visitBase(QualType BaseType, const Record::Base *) {
+    return this->visit(BaseType);
+  }
+  bool visitField(QualType, PrimType, const Record::Field *) {
+    ++Count;
     return true;
   }
+  bool visitFieldComposite(QualType FieldType, const Record::Field *) {
+    return this->visit(FieldType);
+  }
 
-  // Records: fill base classes first, then named fields in declaration
-  // order.
-  if (DestType->isRecordType()) {
-    const Record *R = getRecord(DestType);
-    if (!R)
+private:
+  unsigned Count = 0;
+};
+
+/// Visitor that extracts every scalar element of a source value into
+/// its own local variable.
+template <class Emitter>
+class HLSLElementFlattenVisitor
+    : public Compiler<Emitter>::template HLSLAggregateVisitor<
+          HLSLElementFlattenVisitor<Emitter>> {
+  using VisitorBase = typename Compiler<Emitter>::template HLSLAggregateVisitor<
+      HLSLElementFlattenVisitor<Emitter>>;
+  using HLSLFlatElement = typename Compiler<Emitter>::HLSLFlatElement;
+
+public:
+  HLSLElementFlattenVisitor(Compiler<Emitter> &C, unsigned SrcOffset,
+                            SmallVectorImpl<HLSLFlatElement> &Elements,
+                            unsigned MaxElements, const Expr *E)
+      : VisitorBase(C), CurrentSrcOffset(SrcOffset), Elements(Elements),
+        MaxElements(MaxElements), E(E) {}
+
+  bool isDone() const { return Done; }
+
+  bool visitScalarElem(QualType, PrimType ElemT, unsigned I) {
+    if (checkDone())
+      return false;
+    if (!this->C.emitGetLocal(PT_Ptr, CurrentSrcOffset, E))
       return false;
+    if (!this->C.emitArrayElemPop(ElemT, I, E))
+      return false;
+    return saveToLocal(ElemT);
+  }
 
-    if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(R->getDecl())) {
-      for (const CXXBaseSpecifier &BS : CXXRD->bases()) {
-        const Record::Base *B = R->getBase(BS.getType());
-        assert(B);
-        if (!this->emitGetPtrBase(B->Offset, E))
-          return false;
-        if (!emitHLSLAggregateSplat(SrcT, SrcOffset, BS.getType(), E))
-          return false;
-        if (!this->emitFinishInitPop(E))
-          return false;
-      }
-    }
+  bool visitArrayComposite(QualType ElemType, unsigned I) {
+    if (checkDone())
+      return false;
+    if (!this->C.emitGetLocal(PT_Ptr, CurrentSrcOffset, E))
+      return false;
+    if (!this->C.emitConstUint32(I, E))
+      return false;
+    if (!this->C.emitArrayElemPtrPopUint32(E))
+      return false;
+    return enterSubAggregate(ElemType);
+  }
 
-    for (const Record::Field &F : R->fields()) {
-      if (F.isUnnamedBitField())
-        continue;
+  bool visitBase(QualType BaseType, const Record::Base *B) {
+    if (checkDone())
+      return false;
+    if (!this->C.emitGetLocal(PT_Ptr, CurrentSrcOffset, E))
+      return false;
+    if (!this->C.emitGetPtrBasePop(B->Offset, /*NullOK=*/false, E))
+      return false;
+    return enterSubAggregate(BaseType);
+  }
 
-      QualType FieldType = F.Decl->getType();
-      if (OptPrimType FieldT = classify(FieldType)) {
-        if (!this->emitGetLocal(SrcT, SrcOffset, E))
-          return false;
-        if (!this->emitPrimCast(SrcT, *FieldT, FieldType, E))
-          return false;
-        if (F.isBitField()) {
-          if (!this->emitInitBitField(*FieldT, F.Offset, F.bitWidth(), E))
-            return false;
-        } else {
-          if (!this->emitInitField(*FieldT, F.Offset, E))
-            return false;
-        }
-      } else {
-        if (!this->emitGetPtrField(F.Offset, E))
-          return false;
-        if (!emitHLSLAggregateSplat(SrcT, SrcOffset, FieldType, E))
-          return false;
-        if (!this->emitPopPtr(E))
-          return false;
-      }
+  bool visitField(QualType, PrimType FieldT, const Record::Field *F) {
+    if (checkDone())
+      return false;
+    if (!this->C.emitGetLocal(PT_Ptr, CurrentSrcOffset, E))
+      return false;
+    if (!this->C.emitGetPtrFieldPop(F->Offset, E))
+      return false;
+    if (!this->C.emitLoadPop(FieldT, E))
+      return false;
+    return saveToLocal(FieldT);
+  }
+
+  bool visitFieldComposite(QualType FieldType, const Record::Field *F) {
+    if (checkDone())
+      return false;
+    if (!this->C.emitGetLocal(PT_Ptr, CurrentSrcOffset, E))
+      return false;
+    if (!this->C.emitGetPtrFieldPop(F->Offset, E))
+      return false;
+    return enterSubAggregate(FieldType);
+  }
+
+private:
+  bool checkDone() {
+    if (Done || Elements.size() >= MaxElements) {
+      Done = true;
+      return true;
     }
+    return false;
+  }
+
+  bool saveToLocal(PrimType T) {
+    unsigned Off = this->C.allocateLocalPrimitive(E, T, /*IsConst=*/true);
+    if (!this->C.emitSetLocal(T, Off, E))
+      return false;
+    Elements.push_back({Off, T});
     return true;
   }
 
-  return false;
+  bool enterSubAggregate(QualType SubType) {
+    unsigned Offset =
+        this->C.allocateLocalPrimitive(E, PT_Ptr, /*IsConst=*/true);
+    if (!this->C.emitSetLocal(PT_Ptr, Offset, E))
+      return false;
+    llvm::SaveAndRestore SrcOffsetScope(CurrentSrcOffset, Offset);
+    return this->visit(SubType);
+  }
+
+  unsigned CurrentSrcOffset;
+  SmallVectorImpl<HLSLFlatElement> &Elements;
+  unsigned MaxElements;
+  const Expr *E;
+  bool Done = false;
+};
+
+} // namespace interp
+} // namespace clang
+
+/// Replicate a scalar value into every scalar element of an aggregate.
+/// The scalar is stored in a local at \p SrcOffset and a pointer to the
+/// destination must be on top of the interpreter stack. Each element receives
+/// the scalar, cast to its own type.
+template <class Emitter>
+bool Compiler<Emitter>::emitHLSLAggregateSplat(PrimType SrcT,
+                                               unsigned SrcOffset,
+                                               QualType DestType,
+                                               const Expr *E) {
+  auto ProduceValue = [&](PrimType DestT, QualType DestQT,
+                          const Expr *E) -> bool {
+    if (!this->emitGetLocal(SrcT, SrcOffset, E))
+      return false;
+    return this->emitPrimCast(SrcT, DestT, DestQT, E);
+  };
+  HLSLElementStoreVisitor<Emitter> W(*this, ProduceValue, E);
+  return W.visit(DestType);
 }
 
 /// Return the total number of scalar elements in a type. This is used
@@ -8138,37 +8254,12 @@ bool Compiler<Emitter>::emitHLSLAggregateSplat(PrimType SrcT,
 /// so we never flatten more than the destination can hold.
 template <class Emitter>
 unsigned Compiler<Emitter>::countHLSLFlatElements(QualType Ty) {
-  // Vector and matrix types are treated as flat sequences of elements.
-  if (const auto *VT = Ty->getAs<VectorType>())
-    return VT->getNumElements();
-  if (const auto *MT = Ty->getAs<ConstantMatrixType>())
-    return MT->getNumElementsFlattened();
-  // Arrays: total count is array size * scalar elements per element.
-  if (const auto *AT = Ty->getAsArrayTypeUnsafe()) {
-    const auto *CAT = cast<ConstantArrayType>(AT);
-    return CAT->getZExtSize() * countHLSLFlatElements(CAT->getElementType());
-  }
-  // Records: sum scalar element counts of base classes and named fields.
-  if (Ty->isRecordType()) {
-    const Record *R = getRecord(Ty);
-    if (!R)
-      return 0;
-    unsigned Count = 0;
-    if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(R->getDecl())) {
-      for (const CXXBaseSpecifier &BS : CXXRD->bases())
-        Count += countHLSLFlatElements(BS.getType());
-    }
-    for (const Record::Field &F : R->fields()) {
-      if (F.isUnnamedBitField())
-        continue;
-      Count += countHLSLFlatElements(F.Decl->getType());
-    }
-    return Count;
-  }
-  // Scalar primitive types contribute one element.
   if (canClassify(Ty))
     return 1;
-  return 0;
+  HLSLFlatElementCounter<Emitter> Counter(*this);
+  if (!Counter.visit(Ty))
+    return 0;
+  return Counter.getCount();
 }
 
 /// Walk a source aggregate and extract every scalar element into its own local
@@ -8180,256 +8271,31 @@ bool Compiler<Emitter>::emitHLSLFlattenAggregate(
     QualType SrcType, unsigned SrcOffset,
     SmallVectorImpl<HLSLFlatElement> &Elements, unsigned MaxElements,
     const Expr *E) {
-
-  // Save a scalar value from the stack into a new local and record it.
-  auto saveToLocal = [&](PrimType T) -> bool {
-    unsigned Offset = allocateLocalPrimitive(E, T, /*IsConst=*/true);
-    if (!this->emitSetLocal(T, Offset, E))
-      return false;
-    Elements.push_back({Offset, T});
-    return true;
-  };
-
-  // Save a pointer from the stack into a new local for later use.
-  auto savePtrToLocal = [&]() -> UnsignedOrNone {
-    unsigned Offset = allocateLocalPrimitive(E, PT_Ptr, /*IsConst=*/true);
-    if (!this->emitSetLocal(PT_Ptr, Offset, E))
-      return std::nullopt;
-    return Offset;
-  };
-
-  // Vectors and matrices are flat sequences of elements.
-  unsigned NumElems = 0;
-  QualType ElemType;
-  if (const auto *VT = SrcType->getAs<VectorType>()) {
-    NumElems = VT->getNumElements();
-    ElemType = VT->getElementType();
-  } else if (const auto *MT = SrcType->getAs<ConstantMatrixType>()) {
-    NumElems = MT->getNumElementsFlattened();
-    ElemType = MT->getElementType();
-  }
-  if (NumElems > 0) {
-    PrimType ElemT = classifyPrim(ElemType);
-    for (unsigned I = 0; I != NumElems && Elements.size() < MaxElements; ++I) {
-      if (!this->emitGetLocal(PT_Ptr, SrcOffset, E))
-        return false;
-      if (!this->emitArrayElemPop(ElemT, I, E))
-        return false;
-      if (!saveToLocal(ElemT))
-        return false;
-    }
-    return true;
-  }
-
-  // Arrays: primitive elements are extracted directly; composite elements
-  // require recursion into each sub-aggregate.
-  if (const auto *AT = SrcType->getAsArrayTypeUnsafe()) {
-    const auto *CAT = cast<ConstantArrayType>(AT);
-    QualType ArrElemType = CAT->getElementType();
-    unsigned ArrSize = CAT->getZExtSize();
-
-    if (OptPrimType ElemT = classify(ArrElemType)) {
-      for (unsigned I = 0; I != ArrSize && Elements.size() < MaxElements; ++I) {
-        if (!this->emitGetLocal(PT_Ptr, SrcOffset, E))
-          return false;
-        if (!this->emitArrayElemPop(*ElemT, I, E))
-          return false;
-        if (!saveToLocal(*ElemT))
-          return false;
-      }
-    } else {
-      for (unsigned I = 0; I != ArrSize && Elements.size() < MaxElements; ++I) {
-        if (!this->emitGetLocal(PT_Ptr, SrcOffset, E))
-          return false;
-        if (!this->emitConstUint32(I, E))
-          return false;
-        if (!this->emitArrayElemPtrPopUint32(E))
-          return false;
-        UnsignedOrNone ElemPtrOffset = savePtrToLocal();
-        if (!ElemPtrOffset)
-          return false;
-        if (!emitHLSLFlattenAggregate(ArrElemType, *ElemPtrOffset, Elements,
-                                      MaxElements, E))
-          return false;
-      }
-    }
-    return true;
-  }
-
-  // Records: base classes come first, then named fields in declaration
-  // order.
-  if (SrcType->isRecordType()) {
-    const Record *R = getRecord(SrcType);
-    if (!R)
-      return false;
-
-    if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(R->getDecl())) {
-      for (const CXXBaseSpecifier &BS : CXXRD->bases()) {
-        if (Elements.size() >= MaxElements)
-          break;
-        const Record::Base *B = R->getBase(BS.getType());
-        assert(B);
-        if (!this->emitGetLocal(PT_Ptr, SrcOffset, E))
-          return false;
-        if (!this->emitGetPtrBasePop(B->Offset, /*NullOK=*/false, E))
-          return false;
-        UnsignedOrNone BasePtrOffset = savePtrToLocal();
-        if (!BasePtrOffset)
-          return false;
-        if (!emitHLSLFlattenAggregate(BS.getType(), *BasePtrOffset, Elements,
-                                      MaxElements, E))
-          return false;
-      }
-    }
-
-    for (const Record::Field &F : R->fields()) {
-      if (Elements.size() >= MaxElements)
-        break;
-      if (F.isUnnamedBitField())
-        continue;
-
-      QualType FieldType = F.Decl->getType();
-      if (!this->emitGetLocal(PT_Ptr, SrcOffset, E))
-        return false;
-      if (!this->emitGetPtrFieldPop(F.Offset, E))
-        return false;
-
-      if (OptPrimType FieldT = classify(FieldType)) {
-        if (!this->emitLoadPop(*FieldT, E))
-          return false;
-        if (!saveToLocal(*FieldT))
-          return false;
-      } else {
-        UnsignedOrNone FieldPtrOffset = savePtrToLocal();
-        if (!FieldPtrOffset)
-          return false;
-        if (!emitHLSLFlattenAggregate(FieldType, *FieldPtrOffset, Elements,
-                                      MaxElements, E))
-          return false;
-      }
-    }
-    return true;
-  }
-
-  return false;
+  HLSLElementFlattenVisitor<Emitter> W(*this, SrcOffset, Elements, MaxElements,
+                                       E);
+  bool Result = W.visit(SrcType);
+  return W.isDone() || Result;
 }
 
-/// Populate an HLSL aggregate from a flat list of previously extracted source
-/// elements, casting each to the corresponding destination element type.
+/// Populate an HLSL aggregate from a flat list of previously extracted
+/// source elements, casting each to the corresponding destination element type.
 /// \p ElemIdx tracks the current position in \p Elements and is advanced as
 /// elements are consumed. A pointer to the destination must be on top of the
 /// interpreter stack.
 template <class Emitter>
 bool Compiler<Emitter>::emitHLSLConstructAggregate(
-    QualType DestType, ArrayRef<HLSLFlatElement> Elements, unsigned &ElemIdx,
-    const Expr *E) {
-
-  // Consume the next source element, cast it, and leave it on the stack.
-  auto loadAndCast = [&](PrimType DestT, QualType DestQT) -> bool {
+    QualType DestType, ArrayRef<HLSLFlatElement> Elements, const Expr *E) {
+  unsigned ElemIdx = 0;
+  auto ProduceValue = [&](PrimType DestT, QualType DestQT,
+                          const Expr *E) -> bool {
+    assert(ElemIdx < Elements.size() && "Element index out of bounds");
     const auto &Src = Elements[ElemIdx++];
     if (!this->emitGetLocal(Src.Type, Src.LocalOffset, E))
       return false;
     return this->emitPrimCast(Src.Type, DestT, DestQT, E);
   };
-
-  // Vectors and matrices are flat sequences of elements.
-  unsigned NumElems = 0;
-  QualType ElemType;
-  if (const auto *VT = DestType->getAs<VectorType>()) {
-    NumElems = VT->getNumElements();
-    ElemType = VT->getElementType();
-  } else if (const auto *MT = DestType->getAs<ConstantMatrixType>()) {
-    NumElems = MT->getNumElementsFlattened();
-    ElemType = MT->getElementType();
-  }
-  if (NumElems > 0) {
-    PrimType DestElemT = classifyPrim(ElemType);
-    for (unsigned I = 0; I != NumElems; ++I) {
-      if (!loadAndCast(DestElemT, ElemType))
...
[truncated]

``````````

</details>


https://github.com/llvm/llvm-project/pull/194393


More information about the cfe-commits mailing list