载入中...
搜索中...
未找到
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
17inline bool isQuantDType(DType dt) {
18 return dt == DType::Fp16 || dt == DType::Fp8E4M3 || dt == DType::Fp4E2M1 ||
19 dt == DType::Int8 || dt == DType::Int4;
20}
21
24inline int quantByteSize(DType dt, int count) {
25 switch (dt) {
26 case DType::Fp16: return count * 2;
27 case DType::Fp8E4M3:
28 case DType::Int8: return count;
29 case DType::Fp4E2M1:
30 case DType::Int4: return (count + 1) / 2;
31 default: return count * 4;
32 }
33}
34
35// ---------------------------------------------------------------- fp16
37inline uint16_t f32ToF16(float x) {
38 uint32_t b;
40 std::memcpy(&b, &x, 4);
41 const uint32_t sign = (b >> 16) & 0x8000u;
42 const int32_t e = int32_t((b >> 23) & 0xFFu) - 127 + 15;
43 uint32_t m = b & 0x7FFFFFu;
44 if (e >= 0x1F) return uint16_t(sign | 0x7C00u); // inf / nan -> inf
45 if (e <= 0) {
46 if (e < -10) return uint16_t(sign);
47 m |= 0x800000u;
48 const uint32_t shift = uint32_t(14 - e);
49 uint32_t half = m >> shift;
50 const uint32_t rem = m & ((1u << shift) - 1u);
51 if (rem > (1u << (shift - 1)) || (rem == (1u << (shift - 1)) && (half & 1u))) ++half;
53 return uint16_t(sign | half);
54 }
55 uint32_t half = (uint32_t(e) << 10) | (m >> 13);
56 const uint32_t rem = m & 0x1FFFu;
57 if (rem > 0x1000u || (rem == 0x1000u && (half & 1u))) ++half;
59 return uint16_t(sign | half);
60}
61
63inline float f16ToF32(uint16_t h) {
64 const uint32_t sign = uint32_t(h & 0x8000u) << 16;
65 uint32_t e = (h >> 10) & 0x1Fu;
66 uint32_t m = h & 0x3FFu;
67 uint32_t bits = 0;
68 if (e == 0) {
69 if (m != 0) {
70 e = 127 - 15 + 1;
71 while (!(m & 0x400u)) {
72 m <<= 1;
73 --e;
74 }
75 m &= 0x3FFu;
76 bits = sign | (e << 23) | (m << 13);
77 }
78 } else if (e == 31) {
79 bits = sign | 0x7F800000u | (m << 13);
80 } else {
81 bits = sign | ((e + (127 - 15)) << 23) | (m << 13);
82 }
83 float f;
85 std::memcpy(&f, &bits, 4);
86 return f;
87}
88
89// ---------------------------------------------------------------- fp8 e4m3
91inline float fp8E4M3ToF32(uint8_t v) {
92 const uint32_t sign = uint32_t(v & 0x80u) << 24;
93 const uint32_t e = (v >> 3) & 0xFu;
94 const uint32_t m = v & 0x7u;
95 uint32_t bits;
96 if (e == 0) {
97 bits = m == 0 ? sign : (sign | (uint32_t(127 - 6) << 23) | (m << 20));
98 } else {
99 bits = sign | (uint32_t(e + 120) << 23) | (m << 20); // 2^(e-7) -> fp32 exp e+120
100 }
101 float f;
103 std::memcpy(&f, &bits, 4);
104 return f;
105}
106
107// ---------------------------------------------------------------- fp4 e2m1
109inline float fp4E2M1ToF32(uint8_t nib) {
110 const float sign = (nib & 0x8u) ? -1.f : 1.f;
111 const int e = int((nib >> 1) & 0x3u);
112 const int m = int(nib & 0x1u);
113 const float v = e == 0 ? 0.5f * float(m) : std::ldexp(1.0f, e - 1) * (1.0f + 0.5f * float(m));
114 return sign * v;
115}
116
122inline uint32_t floatToEfm(float x, int expBits, int manBits, int bias) {
124 struct EfmEntry {
125 float value;
126 uint32_t bits;
127 };
128 static std::vector<EfmEntry> *cache = nullptr;
129 static int cacheExp = 0, cacheMan = 0, cacheBias = 0;
130 if (!cache || cacheExp != expBits || cacheMan != manBits || cacheBias != bias) {
131 delete cache;
132 cache = new std::vector<EfmEntry>();
133 const int expCount = 1 << expBits;
134 const int manCount = 1 << manBits;
135 for (uint32_t e = 0; e < uint32_t(expCount); ++e) {
136 for (uint32_t m = 0; m < uint32_t(manCount); ++m) {
137 float v = e == 0
138 ? std::ldexp(float(m), 1 - bias - manBits)
140 : std::ldexp(1.0f + float(m) / float(manCount), int(e) - bias);
141 cache->push_back({v, (e << manBits) | m});
142 }
143 }
145 std::sort(cache->begin(), cache->end(),
146 [](const EfmEntry &a, const EfmEntry &b) { return a.value < b.value; });
147 cacheExp = expBits;
148 cacheMan = manBits;
149 cacheBias = bias;
150 }
151 if (std::isnan(x) || std::isinf(x)) {
152 const uint32_t maxExp = (1u << expBits) - 2u;
153 const uint32_t maxBits = (maxExp << manBits) | ((1u << manBits) - 1u);
154 const uint32_t signBit = 1u << (expBits + manBits);
155 return x < 0 ? (signBit | maxBits) : maxBits;
156 }
157 const float ax = std::fabs(x);
158 size_t lo = 0, hi = cache->size();
159 while (lo + 1 < hi) {
160 const size_t mid = (lo + hi) / 2;
161 if ((*cache)[mid].value <= ax) lo = mid;
162 else hi = mid;
163 }
164 size_t best = lo;
165 if (lo + 1 < cache->size() &&
167 std::fabs((*cache)[lo + 1].value - ax) < std::fabs((*cache)[lo].value - ax))
168 best = lo + 1;
169 const uint32_t bits = (*cache)[best].bits;
170 const uint32_t signBit = 1u << (expBits + manBits);
171 return x < 0 ? (signBit | bits) : bits;
172}
173
175inline uint8_t f32ToFp8E4M3(float x) {
177 return uint8_t(floatToEfm(x, 4, 3, 7));
178}
179
181inline uint8_t f32ToFp4E2M1(float x) {
183 return uint8_t(floatToEfm(x, 2, 1, 1) & 0xFu);
184}
185
188inline float efmMaxMagnitude(int expBits, int manBits, int bias) {
189 const int maxExp = (1 << expBits) - 2;
190 const int maxMan = (1 << manBits) - 1;
192 return std::ldexp(1.0f + float(maxMan) / float(1 << manBits), maxExp - bias);
193}
194
195// ---------------------------------------------------------------- dequant
197inline float dequantValue(DType dt, const uint8_t *bytes, const float *scales, int group,
198 int idx) {
199 switch (dt) {
200 case DType::Fp16: {
201 uint16_t h;
203 std::memcpy(&h, bytes + size_t(idx) * 2, 2);
205 return f16ToF32(h);
206 }
207 case DType::Fp8E4M3:
209 return fp8E4M3ToF32(bytes[static_cast<size_t>(idx)]) *
210 scales[static_cast<size_t>(idx) / size_t(group)];
211 case DType::Int8: {
212 const int8_t v = int8_t(bytes[static_cast<size_t>(idx)]);
214 return float(v) * scales[static_cast<size_t>(idx) / size_t(group)];
215 }
216 case DType::Fp4E2M1: {
217 const uint8_t byte = bytes[static_cast<size_t>(idx) / 2];
218 const uint8_t nib = (idx & 1) ? uint8_t(byte >> 4) : uint8_t(byte & 0xFu);
220 return fp4E2M1ToF32(nib) * scales[static_cast<size_t>(idx) / size_t(group)];
221 }
222 case DType::Int4: {
223 const uint8_t byte = bytes[static_cast<size_t>(idx) / 2];
224 int nib = (idx & 1) ? int(byte >> 4) : int(byte & 0xFu);
225 if (nib >= 8) nib -= 16;
227 return float(nib) * scales[static_cast<size_t>(idx) / size_t(group)];
228 }
229 default: return 0.f;
230 }
231}
232
234inline void dequantizeAll(DType dt, const uint8_t *bytes, const float *scales, int group,
235 int count, float *out) {
236 for (int i = 0; i < count; ++i) out[static_cast<size_t>(i)] = dequantValue(dt, bytes, scales, group, i);
237}
238
239// ---------------------------------------------------------------- quantize
242 std::vector<uint8_t> bytes;
243 std::vector<float> scales;
244 int group = 0;
245};
246
248inline QuantPayload quantize(const float *src, int count, DType dt, int group) {
250 p.group = group <= 0 ? count : group;
251 p.bytes.assign(static_cast<size_t>(quantByteSize(dt, count)), 0);
252 if (dt == DType::Int8 || dt == DType::Int4 || dt == DType::Fp8E4M3 ||
253 dt == DType::Fp4E2M1) {
254 // Per-group block scaling (MXFP-style) for the int and tiny-float
255 // formats: value = dequant(norm) * scale. int formats normalize to
256 // ±maxQ; fp8/fp4 normalize to their largest finite magnitude so the
257 // whole exponent/mantissa range is used (not just [0,1]).
258 const float maxQ = dt == DType::Int8 ? 127.f
259 : dt == DType::Int4 ? 7.f
260 : dt == DType::Fp8E4M3
261 ? efmMaxMagnitude(4, 3, 7)
263 : efmMaxMagnitude(2, 1, 1); // Fp4E2M1
264 const int groups = (count + p.group - 1) / p.group;
265 p.scales.resize(static_cast<size_t>(groups));
266 for (int g = 0; g < groups; ++g) {
267 float maxAbs = 0.f;
268 const int begin = g * p.group;
269 const int end = std::min(count, begin + p.group);
270 for (int i = begin; i < end; ++i) maxAbs = std::max(maxAbs, std::fabs(src[static_cast<size_t>(i)]));
271 float scale = maxAbs / float(maxQ);
272 if (scale == 0.f) scale = 1.f;
273 p.scales[static_cast<size_t>(g)] = scale;
274 for (int i = begin; i < end; ++i) {
275 const float norm = src[static_cast<size_t>(i)] / scale;
276 if (dt == DType::Fp8E4M3) {
277 p.bytes[static_cast<size_t>(i)] = f32ToFp8E4M3(norm);
278 continue;
279 }
280 if (dt == DType::Fp4E2M1) {
281 const uint8_t nib = f32ToFp4E2M1(norm);
282 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
283 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
284 else byte = uint8_t((byte & 0xF0u) | nib);
285 continue;
286 }
287 int qv = int(std::floor(norm + 0.5f));
288 const int qBound = int(maxQ);
289 qv = std::max(-qBound - 1, std::min(qBound, qv));
290 if (dt == DType::Int8) {
291 p.bytes[static_cast<size_t>(i)] = uint8_t(int8_t(qv));
292 } else {
293 const int nib = qv & 0xFu;
294 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
295 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
296 else byte = uint8_t((byte & 0xF0u) | uint8_t(nib));
297 }
298 }
299 }
300 return p;
301 }
302 for (int i = 0; i < count; ++i) {
303 const float v = src[static_cast<size_t>(i)];
304 switch (dt) {
305 case DType::Fp16: {
306 const uint16_t h = f32ToF16(v);
308 std::memcpy(p.bytes.data() + size_t(i) * 2, &h, 2);
309 break;
310 }
311 case DType::Fp8E4M3: p.bytes[static_cast<size_t>(i)] = f32ToFp8E4M3(v); break;
312 case DType::Fp4E2M1: {
313 const uint8_t nib = f32ToFp4E2M1(v);
314 uint8_t &byte = p.bytes[static_cast<size_t>(i) / 2];
315 if (i & 1) byte = uint8_t((byte & 0x0Fu) | uint8_t(nib << 4));
316 else byte = uint8_t((byte & 0xF0u) | nib);
317 break;
318 }
319 default: break;
320 }
321 }
322 return p;
323}
324
325} // namespace q
326} // namespace eve::tensor
327
328#endif // EVE_TENSOR_QUANT_H
double value
float x
Definition AnimClip.cpp:738
int mid
Definition AnimSmr.cpp:120
building::EdgeCurveGroup group
int ax
Definition CaveMesh.cpp:113
glm::vec4 p[6]
tensor::Graph g
Definition GpuGraph.cpp:7
std::array< double, 10 > q
float v
int h
std::array< float, 3 > scale
std::uint64_t bytes
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
uint32_t groups
Definition OnnxGpgpu.cpp:39
std::vector< float > scales
Definition OnnxLstm.cpp:27
int idx
float f
float begin
float bias
std::uint32_t count
std::map< Cell, int > best
float m[16]
void dequantizeAll(DType dt, const uint8_t *bytes, const float *scales, int group, int count, float *out)
Dequantize all.
Definition Quant.h:234
float fp4E2M1ToF32(uint8_t nib)
Fp 4 e 2 m 1 to f 32.
Definition Quant.h:109
QuantPayload quantize(const float *src, int count, DType dt, int group)
Quantize.
Definition Quant.h:248
float f16ToF32(uint16_t h)
F 16 to f 32.
Definition Quant.h:63
uint32_t floatToEfm(float x, int expBits, int manBits, int bias)
Float to efm.
Definition Quant.h:122
uint8_t f32ToFp8E4M3(float x)
F 32 to fp 8 e 4 m 3.
Definition Quant.h:175
float efmMaxMagnitude(int expBits, int manBits, int bias)
Efm max magnitude.
Definition Quant.h:188
float dequantValue(DType dt, const uint8_t *bytes, const float *scales, int group, int idx)
Dequant value.
Definition Quant.h:197
uint8_t f32ToFp4E2M1(float x)
F 32 to fp 4 e 2 m 1.
Definition Quant.h:181
int quantByteSize(DType dt, int count)
Quant byte size.
Definition Quant.h:24
float fp8E4M3ToF32(uint8_t v)
Fp 8 e 4 m 3 to f 32.
Definition Quant.h:91
bool isQuantDType(DType dt)
True when quant d type.
Definition Quant.h:17
uint16_t f32ToF16(float x)
F 32 to f 16.
Definition Quant.h:37
DType
Tensor element types.
Definition Tensor.h:24
QuantPayload public API.
Definition Quant.h:241
std::vector< float > scales
Definition Quant.h:243
std::vector< uint8_t > bytes
Definition Quant.h:242