载入中...
搜索中...
未找到
Quant.h
浏览该文件的文档.
1#ifndef EVE_TENSOR_QUANT_H
2#define EVE_TENSOR_QUANT_H
3
4#include "tensor/Tensor.h"
5
6#include <algorithm>
7#include <cmath>
8#include <cstdint>
9#include <cstring>
10#include <vector>
11
12namespace eve::tensor {
13namespace q {
14
16inline bool isQuantDType(DType dt) {
17 return dt == DType::Fp16 || dt == DType::Fp8E4M3 || dt == DType::Fp4E2M1 ||
18 dt == DType::Int8 || dt == DType::Int4;
19}
20
22inline int quantByteSize(DType dt, int count) {
23 switch (dt) {
24 case DType::Fp16: return count * 2;
25 case DType::Fp8E4M3:
26 case DType::Int8: return count;
27 case DType::Fp4E2M1:
28 case DType::Int4: return (count + 1) / 2;
29 default: return count * 4;
30 }
31}
32
33// ---------------------------------------------------------------- fp16
34inline uint16_t f32ToF16(float x) {
35 uint32_t b;
36 std::memcpy(&b, &x, 4);
37 const uint32_t sign = (b >> 16) & 0x8000u;
38 const int32_t e = int32_t((b >> 23) & 0xFFu) - 127 + 15;
39 uint32_t m = b & 0x7FFFFFu;
40 if (e >= 0x1F) return uint16_t(sign | 0x7C00u); // inf / nan -> inf
41 if (e <= 0) {
42 if (e < -10) return uint16_t(sign);
43 m |= 0x800000u;
44 const uint32_t shift = uint32_t(14 - e);
45 uint32_t half = m >> shift;
46 const uint32_t rem = m & ((1u << shift) - 1u);
47 if (rem > (1u << (shift - 1)) || (rem == (1u << (shift - 1)) && (half & 1u))) ++half;
48 return uint16_t(sign | half);
49 }
50 uint32_t half = (uint32_t(e) << 10) | (m >> 13);
51 const uint32_t rem = m & 0x1FFFu;
52 if (rem > 0x1000u || (rem == 0x1000u && (half & 1u))) ++half;
53 return uint16_t(sign | half);
54}
55
56inline float f16ToF32(uint16_t h) {
57 const uint32_t sign = uint32_t(h & 0x8000u) << 16;
58 uint32_t e = (h >> 10) & 0x1Fu;
59 uint32_t m = h & 0x3FFu;
60 uint32_t bits = 0;
61 if (e == 0) {
62 if (m != 0) {
63 e = 127 - 15 + 1;
64 while (!(m & 0x400u)) {
65 m <<= 1;
66 --e;
67 }
68 m &= 0x3FFu;
69 bits = sign | (e << 23) | (m << 13);
70 }
71 } else if (e == 31) {
72 bits = sign | 0x7F800000u | (m << 13);
73 } else {
74 bits = sign | ((e + (127 - 15)) << 23) | (m << 13);
75 }
76 float f;
77 std::memcpy(&f, &bits, 4);
78 return f;
79}
80
81// ---------------------------------------------------------------- fp8 e4m3
82inline float fp8E4M3ToF32(uint8_t v) {
83 const uint32_t sign = uint32_t(v & 0x80u) << 24;
84 const uint32_t e = (v >> 3) & 0xFu;
85 const uint32_t m = v & 0x7u;
86 uint32_t bits;
87 if (e == 0) {
88 bits = m == 0 ? sign : (sign | (uint32_t(127 - 6) << 23) | (m << 20));
89 } else {
90 bits = sign | (uint32_t(e + 120) << 23) | (m << 20); // 2^(e-7) -> fp32 exp e+120
91 }
92 float f;
93 std::memcpy(&f, &bits, 4);
94 return f;
95}
96
97// ---------------------------------------------------------------- fp4 e2m1
98inline float fp4E2M1ToF32(uint8_t nib) {
99 const float sign = (nib & 0x8u) ? -1.f : 1.f;
100 const int e = int((nib >> 1) & 0x3u);
101 const int m = int(nib & 0x1u);
102 const float v = e == 0 ? 0.5f * float(m) : std::ldexp(1.0f, e - 1) * (1.0f + 0.5f * float(m));
103 return sign * v;
104}
105
110inline uint32_t floatToEfm(float x, int expBits, int manBits, int bias) {
111 struct EfmEntry {
112 float value;
113 uint32_t bits;
114 };
115 static std::vector<EfmEntry> *cache = nullptr;
116 static int cacheExp = 0, cacheMan = 0, cacheBias = 0;
117 if (!cache || cacheExp != expBits || cacheMan != manBits || cacheBias != bias) {
118 delete cache;
119 cache = new std::vector<EfmEntry>();
120 const int expCount = 1 << expBits;
121 const int manCount = 1 << manBits;
122 for (uint32_t e = 0; e < uint32_t(expCount); ++e) {
123 for (uint32_t m = 0; m < uint32_t(manCount); ++m) {
124 float v = e == 0
125 ? std::ldexp(float(m), 1 - bias - manBits)
126 : std::ldexp(1.0f + float(m) / float(manCount), int(e) - bias);
127 cache->push_back({v, (e << manBits) | m});
128 }
129 }
130 std::sort(cache->begin(), cache->end(),
131 [](const EfmEntry &a, const EfmEntry &b) { return a.value < b.value; });
132 cacheExp = expBits;
133 cacheMan = manBits;
134 cacheBias = bias;
135 }
136 if (std::isnan(x) || std::isinf(x)) {
137 const uint32_t maxExp = (1u << expBits) - 2u;
138 const uint32_t maxBits = (maxExp << manBits) | ((1u << manBits) - 1u);
139 const uint32_t signBit = 1u << (expBits + manBits);
140 return x < 0 ? (signBit | maxBits) : maxBits;
141 }
142 const float ax = std::fabs(x);
143 size_t lo = 0, hi = cache->size();
144 while (lo + 1 < hi) {
145 const size_t mid = (lo + hi) / 2;
146 if ((*cache)[mid].value <= ax) lo = mid;
147 else hi = mid;
148 }
149 size_t best = lo;
150 if (lo + 1 < cache->size() &&
151 std::fabs((*cache)[lo + 1].value - ax) < std::fabs((*cache)[lo].value - ax))
152 best = lo + 1;
153 const uint32_t bits = (*cache)[best].bits;
154 const uint32_t signBit = 1u << (expBits + manBits);
155 return x < 0 ? (signBit | bits) : bits;
156}
157
158inline uint8_t f32ToFp8E4M3(float x) {
159 return uint8_t(floatToEfm(x, 4, 3, 7));
160}
161
162inline uint8_t f32ToFp4E2M1(float x) {
163 return uint8_t(floatToEfm(x, 2, 1, 1) & 0xFu);
164}
165
167inline float efmMaxMagnitude(int expBits, int manBits, int bias) {
168 const int maxExp = (1 << expBits) - 2;
169 const int maxMan = (1 << manBits) - 1;
170 return std::ldexp(1.0f + float(maxMan) / float(1 << manBits), maxExp - bias);
171}
172
173// ---------------------------------------------------------------- dequant
174inline float dequantValue(DType dt, const uint8_t *bytes, const float *scales, int group,
175 int idx) {
176 switch (dt) {
177 case DType::Fp16: {
178 uint16_t h;
179 std::memcpy(&h, bytes + size_t(idx) * 2, 2);
180 return f16ToF32(h);
181 }
182 case DType::Fp8E4M3:
183 return fp8E4M3ToF32(bytes[static_cast<size_t>(idx)]) *
184 scales[static_cast<size_t>(idx) / size_t(group)];
185 case DType::Int8: {
186 const int8_t v = int8_t(bytes[static_cast<size_t>(idx)]);
187 return float(v) * scales[static_cast<size_t>(idx) / size_t(group)];
188 }
189 case DType::Fp4E2M1: {
190 const uint8_t byte = bytes[static_cast<size_t>(idx) / 2];
191 const uint8_t nib = (idx & 1) ? uint8_t(byte >> 4) : uint8_t(byte & 0xFu);
192 return fp4E2M1ToF32(nib) * scales[static_cast<size_t>(idx) / size_t(group)];
193 }
194 case DType::Int4: {
195 const uint8_t byte = bytes[static_cast<size_t>(idx) / 2];
196 int nib = (idx & 1) ? int(byte >> 4) : int(byte & 0xFu);
197 if (nib >= 8) nib -= 16;
198 return float(nib) * scales[static_cast<size_t>(idx) / size_t(group)];
199 }
200 default: return 0.f;
201 }
202}
203
204inline void dequantizeAll(DType dt, const uint8_t *bytes, const float *scales, int group,
205 int count, float *out) {
206 for (int i = 0; i < count; ++i) out[static_cast<size_t>(i)] = dequantValue(dt, bytes, scales, group, i);
207}
208
209// ---------------------------------------------------------------- quantize
211 std::vector<uint8_t> bytes;
212 std::vector<float> scales;
213 int group = 0;
214};
215
216inline QuantPayload quantize(const float *src, int count, DType dt, int group) {
218 p.group = group <= 0 ? count : group;
219 p.bytes.assign(static_cast<size_t>(quantByteSize(dt, count)), 0);
220 if (dt == DType::Int8 || dt == DType::Int4 || dt == DType::Fp8E4M3 ||
221 dt == DType::Fp4E2M1) {
222 // Per-group block scaling (MXFP-style) for the int and tiny-float
223 // formats: value = dequant(norm) * scale. int formats normalize to
224 // ±maxQ; fp8/fp4 normalize to their largest finite magnitude so the
225 // whole exponent/mantissa range is used (not just [0,1]).
226 const float maxQ = dt == DType::Int8 ? 127.f
227 : dt == DType::Int4 ? 7.f
228 : dt == DType::Fp8E4M3
229 ? efmMaxMagnitude(4, 3, 7)
230 : efmMaxMagnitude(2, 1, 1); // Fp4E2M1
231 const int groups = (count + p.group - 1) / p.group;
232 p.scales.resize(static_cast<size_t>(groups));
233 for (int g = 0; g < groups; ++g) {
234 float maxAbs = 0.f;
235 const int begin = g * p.group;
236 const int end = std::min(count, begin + p.group);
237 for (int i = begin; i < end; ++i) maxAbs = std::max(maxAbs, std::fabs(src[static_cast<size_t>(i)]));
238 float scale = maxAbs / float(maxQ);
239 if (scale == 0.f) scale = 1.f;
240 p.scales[static_cast<size_t>(g)] = scale;
241 for (int i = begin; i < end; ++i) {
242 const float norm = src[static_cast<size_t>(i)] / scale;
243 if (dt == DType::Fp8E4M3) {
244 p.bytes[static_cast<size_t>(i)] = f32ToFp8E4M3(norm);
245 continue;
246 }
247 if (dt == DType::Fp4E2M1) {
248 const uint8_t nib = f32ToFp4E2M1(norm);
249 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
250 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
251 else byte = uint8_t((byte & 0xF0u) | nib);
252 continue;
253 }
254 int qv = int(std::floor(norm + 0.5f));
255 const int qBound = int(maxQ);
256 qv = std::max(-qBound - 1, std::min(qBound, qv));
257 if (dt == DType::Int8) {
258 p.bytes[static_cast<size_t>(i)] = uint8_t(int8_t(qv));
259 } else {
260 const int nib = qv & 0xFu;
261 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
262 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
263 else byte = uint8_t((byte & 0xF0u) | uint8_t(nib));
264 }
265 }
266 }
267 return p;
268 }
269 for (int i = 0; i < count; ++i) {
270 const float v = src[static_cast<size_t>(i)];
271 switch (dt) {
272 case DType::Fp16: {
273 const uint16_t h = f32ToF16(v);
274 std::memcpy(p.bytes.data() + size_t(i) * 2, &h, 2);
275 break;
276 }
277 case DType::Fp8E4M3: p.bytes[static_cast<size_t>(i)] = f32ToFp8E4M3(v); break;
278 case DType::Fp4E2M1: {
279 const uint8_t nib = f32ToFp4E2M1(v);
280 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
281 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
282 else byte = uint8_t((byte & 0xF0u) | nib);
283 break;
284 }
285 default: break;
286 }
287 }
288 return p;
289}
290
291} // namespace q
292} // namespace eve::tensor
293
294#endif // EVE_TENSOR_QUANT_H
std::string value
int x
Definition Grass.cpp:135
int h
const FusedGroup & group
uint32_t a
uint32_t b
int idx
float f
glm::vec4 p[6]
int v
float scale
Definition TreeMesh.cpp:122
float m[16]
void dequantizeAll(DType dt, const uint8_t *bytes, const float *scales, int group, int count, float *out)
Definition Quant.h:204
float fp4E2M1ToF32(uint8_t nib)
Definition Quant.h:98
QuantPayload quantize(const float *src, int count, DType dt, int group)
Definition Quant.h:216
float f16ToF32(uint16_t h)
Definition Quant.h:56
uint32_t floatToEfm(float x, int expBits, int manBits, int bias)
Definition Quant.h:110
uint8_t f32ToFp8E4M3(float x)
Definition Quant.h:158
float efmMaxMagnitude(int expBits, int manBits, int bias)
Definition Quant.h:167
float dequantValue(DType dt, const uint8_t *bytes, const float *scales, int group, int idx)
Definition Quant.h:174
uint8_t f32ToFp4E2M1(float x)
Definition Quant.h:162
int quantByteSize(DType dt, int count)
Definition Quant.h:22
float fp8E4M3ToF32(uint8_t v)
Definition Quant.h:82
bool isQuantDType(DType dt)
Definition Quant.h:16
uint16_t f32ToF16(float x)
Definition Quant.h:34
DType
Tensor element types.
Definition Tensor.h:22
std::vector< float > scales
Definition Quant.h:212
std::vector< uint8_t > bytes
Definition Quant.h:211