#pragma once
#include <cstdint>
#include <cstdlib>
namespace vt {
enum class FOp : uint8_t {
kAdd, kMul, kSilu, kSigmoid, kRmsNorm, kSiluMul, kSigmoidGate, kRmsNormGated, kRope, kQuantFp8, kQuantFp4, kAttnQkNormRopeGate, };
enum class FReduce : uint8_t {
kNone,
kMeanSquare,
};
enum class FKind : uint8_t {
kUnused = 0,
kRow,
kWeight,
kResidual,
kAux,
};
constexpr uint8_t kNoOperand = 0xFF;
constexpr int kNoFastOp = -1;
constexpr int kMaxFusedSteps = 8;
constexpr int kMaxFusedOperands = 8;
constexpr int kMaxStepIns = 3;
struct FOperandSlot {
FKind kind = FKind::kUnused;
const char* name = nullptr; };
struct FStep {
FOp op = FOp::kAdd;
uint8_t out = 0; uint8_t in[kMaxStepIns] = {0, 0, 0}; uint8_t nin = 0; uint8_t out2 = kNoOperand; FReduce reduce = FReduce::kNone;
bool gemma = false; bool sigmoid_gate = false; bool norm_full_width = false;
};
struct FusedRecipe {
FStep steps[kMaxFusedSteps] = {};
FOperandSlot operands[kMaxFusedOperands] = {};
int n = 0; int n_operands = 0; const char* name = nullptr;
int fast_op = kNoFastOp; };
inline int FusedTier() {
const char* e = std::getenv("VT_FUSED_TIER");
return (e != nullptr && e[0] == '1') ? 1 : 0;
}
inline bool RecipeIsTier1Able(const FusedRecipe& r) {
for (int s = 0; s < r.n; ++s) {
switch (r.steps[s].op) {
case FOp::kAdd:
case FOp::kMul:
case FOp::kSilu:
case FOp::kSigmoid:
case FOp::kRmsNorm:
break;
default:
return false;
}
}
return true;
}
}