Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions c/Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -526,6 +526,9 @@ iobench$(EXE): iobench.c compat.h
tests/test_serve_sentinel$(EXE): tests/test_serve_sentinel.c compat.h
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)

tests/test_ue8m0$(EXE): tests/test_ue8m0.c st.h json.h compat.h
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)

tests/test_json$(EXE): tests/test_json.c json.h
$(CC) $(CFLAGS) $< -o $@ $(LDFLAGS)

Expand Down
17 changes: 14 additions & 3 deletions c/colibri.c
Original file line number Diff line number Diff line change
Expand Up @@ -1672,11 +1672,22 @@ static void qt_from_disk(Model *m, const char *name, int O, int I, int bits, int
* block scales; everything else is per-row O. Using the per-row bound for a
* grouped/blocked format would reject a legitimate container (fmt=5 regressed
* exactly that way). */
st_read_f32_cap(&m->S,sn,t->s,
/* fmt=8 goes through st_read_scale_f32: the block-scale GEOMETRY is the
* same in a container we repacked (f32 scales) and in a native fp8
* checkpoint (UE8M0, one byte per block) -- same shape, same meaning,
* same multiply. Only the encoding of the number differs. The reader
* accepts either and always yields f32, so matmul_fp8 stays a single
* implementation with no branch in the hot loop.
*
* Every OTHER format still goes through st_read_f32_cap exactly as
* before: fmt 0/1/2/4/5/6 are byte-for-byte unchanged, and an f32-scaled
* fmt=8 container behaves identically too (st_read_scale_f32 dispatches
* to st_read_f32 for dtype F32, the same call it made before). */
if(fmt==8) st_read_scale_f32(&m->S,sn,t->s,fp8_nblk(O)*fp8_nblk(I),drop);
else st_read_f32_cap(&m->S,sn,t->s,
fmt==4 ? (int64_t)O*((I+gs-1)/gs) :
fmt==5 ? (int64_t)O*i3_groups(I) :
fmt==6 ? (int64_t)1 :
fmt==8 ? fp8_nblk(O)*fp8_nblk(I) : (int64_t)O, drop);
fmt==6 ? (int64_t)1 : (int64_t)O, drop);
} else {
if(!t->qf && !t->q8 && !t->q4) qt_alloc(t,O,I,bits);
if(t->fmt==0) st_read_f32_cap(&m->S,name,t->qf,(int64_t)O*I,drop);
Expand Down
109 changes: 105 additions & 4 deletions c/st.h
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <math.h> /* ldexpf per ue8m0_to_f32 */
#include <stdint.h>
#include <fcntl.h>
#include <errno.h>
Expand All @@ -29,7 +30,7 @@ typedef struct {
int fd;
int64_t off; /* offset assoluto del dato dentro al file */
int64_t nbytes;
int dtype; /* 0=BF16 1=F16 2=F32 */
int dtype; /* 0=BF16 1=F16 2=F32 3=U8/I8 4=F8_E4M3 5=F8_E8M0 6=I64 */
int64_t numel;
} st_tensor;

Expand Down Expand Up @@ -78,9 +79,40 @@ static int st_dtype_code(const char *s) {
if (!strcmp(s, "F32")) return 2;
if (!strcmp(s, "U8")) return 3; /* dati quantizzati (int4 packed / int8) */
if (!strcmp(s, "I8")) return 3;
/* --- tipi dei checkpoint nativi fp8 (DeepSeek-V4, GLM-5.2-FP8 non ripacchettati) ---
* PRIMA di questi, st_init faceva exit(1) su un checkpoint DeepSeek-V4 al primo
* tensore I64, senza arrivare ai pesi. Sono INDICIZZATI qui e letti dal percorso
* dei byte grezzi (st_read_raw); i lettori float li RIFIUTANO PER NOME invece di
* caderci dentro -- vedi il commento in st_read_f32. Il loro codice numerico e'
* nuovo e nessun codice esistente cambia: 0/1/2/3 restano quelli di prima. */
if (!strcmp(s, "F8_E4M3") || !strcmp(s, "F8_E4M3FN") ||
!strcmp(s, "float8_e4m3fn")) return 4;
if (!strcmp(s, "F8_E8M0") || !strcmp(s, "F8_E8M0FNU")) return 5;
if (!strcmp(s, "I64") || !strcmp(s, "U64")) return 6;
fprintf(stderr, "unsupported dtype: %s\n", s); exit(1);
}

/* Byte per elemento. UNICO posto che lo sa: prima la formula era ripetuta in tre
* punti come `dtype==2 ? 4 : 2`, che con soli 0/1/2/3 era corretta e con i tipi
* nuovi avrebbe silenziosamente detto "2 byte" per un I64 da 8. */
static inline int st_dtype_esz(int dtype) {
switch (dtype) {
case 2: return 4; /* F32 */
case 3: case 4: case 5: return 1; /* U8/I8, F8_E4M3, F8_E8M0 */
case 6: return 8; /* I64/U64 */
default: return 2; /* BF16, F16 */
}
}

/* Nome leggibile, per i messaggi di rifiuto. */
static inline const char *st_dtype_name(int dtype) {
switch (dtype) {
case 0: return "BF16"; case 1: return "F16"; case 2: return "F32";
case 3: return "U8/I8"; case 4: return "F8_E4M3"; case 5: return "F8_E8M0";
case 6: return "I64"; default: return "?";
}
}

static inline float bf16_to_f32(uint16_t h) {
uint32_t u = (uint32_t)h << 16; float f; memcpy(&f, &u, 4); return f;
}
Expand Down Expand Up @@ -506,7 +538,7 @@ static void st_init_multi(shards *S, const char *snap_dir, const char *extra_dir
* into a caller-sized buffer, so a header with numel != nbytes/esz is an
* OOB write primitive. U8/I8 (raw quant bytes) are read by byte count, so
* their numel is unused by the read path and legitimately may differ. */
{ int esz = t->dtype==2 ? 4 : (t->dtype==3 ? 1 : 2);
{ int esz = st_dtype_esz(t->dtype);
if (t->dtype != 3 && t->nbytes != numel * (int64_t)esz) {
fprintf(stderr, "%s: tensor '%s' numel %lld disagrees with byte span %lld (esz %d)\n",
files[fi], name, (long long)numel, (long long)t->nbytes, esz); exit(1); } }
Expand Down Expand Up @@ -637,7 +669,15 @@ static int64_t st_read_f32(shards *S, const char *name, float *out, int drop) {
* (numel elementi da un raw di soli nbytes) sforano il buffer del chiamante,
* che e' dimensionato sul config, non sul file. Il chiamante che alloca su
* st_numel resta coerente; questo blocca l'ingresso ostile a monte. */
int esz = (t->dtype == 2) ? 4 : 2;
/* I tipi non-float si leggono con st_read_raw, non qui. Senza questo rifiuto
* cadrebbero nel ramo `else` in fondo, che assume F16: un tensore F8_E4M3 o
* I64 verrebbe letto come mezze-precisioni e produrrebbe numeri plausibili e
* sbagliati, in silenzio. Con soli i dtype 0/1/2/3 il fallthrough era corretto;
* dal momento in cui ne esistono altri, non lo e' piu'. */
if (t->dtype >= 3) {
fprintf(stderr, "%s: tensor '%s' is %s — not a float tensor; read it with st_read_raw\n",
name, name, st_dtype_name(t->dtype)); exit(1); }
int esz = st_dtype_esz(t->dtype);
if (t->numel < 0 || t->numel > t->nbytes / esz || t->numel * (int64_t)esz != t->nbytes) {
fprintf(stderr, "%s: tensor '%s' shape/bytes mismatch (numel %lld, %lld bytes, dtype %d) — refusing (hostile or corrupt file)\n",
name, name, (long long)t->numel, (long long)t->nbytes, t->dtype); exit(1); }
Expand Down Expand Up @@ -677,6 +717,64 @@ static int64_t st_nbytes(shards *S, const char *name) {
st_tensor *t = st_find(S, name); return t ? t->nbytes : -1;
}

/* --- ue8m0_to_f32 / st_read_scale_f32: sidecar di scale a 1 byte ------------
*
* UE8M0 e' un esponente potenza-di-due senza segno e senza mantissa: il valore
* e' 2^(v-127), e 0xff e' NaN. Un byte per scala invece di quattro.
*
* Serve perche' un checkpoint fp8 nativo (DeepSeek-V4, e in generale
* quantization_config.scale_fmt == "ue8m0") scrive le scale di blocco cosi',
* mentre i container ripacchettati da noi le scrivono in f32. La GEOMETRIA e'
* identica -- stessa forma, stesso significato, stessa moltiplicazione -- cambia
* solo la codifica del numero.
*
* Si espande a f32 UNA VOLTA al caricamento invece di decodificare nel kernel:
* le scale sono ~1/16384 dei byte dei pesi (mezzo MB per gli 8,4 GB densi di
* DeepSeek-V4), quindi il costo in memoria e' trascurabile e matmul_fp8 resta
* UNA sola implementazione, senza un ramo dentro il ciclo caldo. E' la stessa
* scelta gia' fatta in kimi_k3.c per le scale ue8m0 di MXFP4 (`mx4_scale`).
*
* NaN: 0xff decodifica a un NaN IEEE reale e lo si lascia propagare, coerente
* con la politica gia' documentata per i pesi fp8 in quant.h -- la rete di
* sicurezza sta a valle, sul sampler (test_logit_nan.c), non qui. */
static inline float ue8m0_to_f32(uint8_t v) {
if (v == 0xff) { uint32_t n = 0x7fc00000u; float f; memcpy(&f, &n, 4); return f; }
/* ldexpf e NON il trucco `(uint32_t)v << 23`: quel trucco e' esatto per
* v in [1,254], ma a v==0 produce il pattern di bit 0x00000000, che in IEEE
* 754 e' ZERO ESATTO e non 2^-127. Un blocco di pesi con quella scala
* verrebbe azzerato invece che reso quasi-zero -- differenza piccola in
* ampiezza, ma e' comunque un valore sbagliato, e 2^-127 e' rappresentabile
* come subnormale. Costa solo al caricamento (una volta per scala), quindi
* si paga la chiamata e si tiene la correttezza. */
return ldexpf(1.0f, (int)v - 127);
}

/* Legge un sidecar di scale in `out` come f32, accettando SIA F32 SIA F8_E8M0.
* Rifiuta qualunque altro dtype per nome. `cap` e' il numero massimo di float
* che il chiamante ha allocato, come in st_read_f32_cap. */
static int64_t st_read_scale_f32(shards *S, const char *name, float *out, int64_t cap, int drop) {
st_tensor *t = st_find(S, name);
if (!t) { fprintf(stderr, "missing tensor: %s\n", name); exit(1); }
if (t->numel < 0 || t->numel > cap) {
fprintf(stderr, "scale %s: numel %lld exceeds destination capacity %lld\n",
name, (long long)t->numel, (long long)cap); exit(1); }
if (t->dtype == 2 || t->dtype == 0 || t->dtype == 1) return st_read_f32(S, name, out, drop);
if (t->dtype != 5) {
fprintf(stderr, "scale %s: dtype %s is neither F32 nor F8_E8M0\n",
name, st_dtype_name(t->dtype)); exit(1); }
/* stessa validazione byte-vs-numel dei percorsi float: 1 byte per scala */
if (t->nbytes != t->numel) {
fprintf(stderr, "scale %s: F8_E8M0 numel %lld disagrees with %lld bytes\n",
name, (long long)t->numel, (long long)t->nbytes); exit(1); }
uint8_t *raw = (uint8_t*)malloc((size_t)t->nbytes);
if (!raw) { fprintf(stderr, "malloc %lld bytes for scale %s failed\n", (long long)t->nbytes, name); exit(1); }
st_pread_full(t->fd, raw, t->nbytes, t->off, "pread ue8m0 scale");
for (int64_t i = 0; i < t->numel; i++) out[i] = ue8m0_to_f32(raw[i]);
free(raw);
if (drop) posix_fadvise(t->fd, t->off, t->nbytes, POSIX_FADV_DONTNEED);
return t->numel;
}

/* legge i byte GREZZI di un tensore (nessuna conversione di dtype): per i pesi gia'
* quantizzati int4/int8 del nostro container (dtype U8). drop=1 -> fadvise DONTNEED. */
static void st_read_raw(shards *S, const char *name, void *out, int drop) {
Expand All @@ -692,7 +790,10 @@ static void st_read_raw(shards *S, const char *name, void *out, int drop) {
static void st_read_slice_f32(shards *S, const char *name, int64_t elem_off, int64_t n_elems, float *out, int drop) {
st_tensor *t = st_find(S, name);
if (!t) { fprintf(stderr, "missing tensor: %s\n", name); exit(1); }
int esz = (t->dtype == 2) ? 4 : 2;
if (t->dtype >= 3) { /* stesso motivo di st_read_f32 sopra */
fprintf(stderr, "slice %s: tensor is %s — not a float tensor\n",
name, st_dtype_name(t->dtype)); exit(1); }
int esz = st_dtype_esz(t->dtype);
if (elem_off < 0 || n_elems < 0 || elem_off > t->numel || n_elems > t->numel - elem_off) { /* keep the slice inside the tensor; subtraction avoids overflow (#1) */
fprintf(stderr, "slice %s [%lld,+%lld) out of tensor bounds (numel %lld)\n",
name, (long long)elem_off, (long long)n_elems, (long long)t->numel); exit(1); }
Expand Down
91 changes: 91 additions & 0 deletions c/tests/test_ue8m0.c
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
/* test_ue8m0 — the UE8M0 block-scale decode, and the dtype table it arrived with.
*
* UE8M0 is an unsigned power-of-two exponent with no mantissa: the value is
* 2^(v-127), and 0xff is NaN. Native fp8 checkpoints (DeepSeek-V4, and anything
* with quantization_config.scale_fmt == "ue8m0") write block scales this way
* instead of as f32.
*
* The whole domain is 256 values, so this checks all of them rather than
* sampling. That is not thoroughness for its own sake: the first implementation
* used the obvious bit trick, `(uint32_t)v << 23`, which is exact for
* v in [1,254] and silently wrong at v == 0 -- the all-zero bit pattern is EXACT
* ZERO in IEEE 754, not 2^-127. A weight block carrying that scale would have
* been zeroed rather than made almost-zero. Sampling would not have found it;
* the boundary is the whole bug.
*
* Also pins st_dtype_esz, because that function replaced three copies of a
* `dtype==2 ? 4 : 2` ternary. That ternary was correct while only four dtypes
* existed and would have claimed 2 bytes for an 8-byte I64 the moment a fifth
* appeared. A wrong element size is an out-of-bounds read, not a wrong number.
*/
#include <stdio.h>
#include <math.h>
#include <string.h>
#include "../st.h"

static int fails = 0;

static void check(int cond, const char *what)
{
if (!cond) { printf(" FAIL %s\n", what); fails++; }
}

int main(void)
{
/* --- the whole UE8M0 domain --- */
int bad = 0;
for (int v = 0; v < 256; v++) {
float got = ue8m0_to_f32((uint8_t)v);
if (v == 0xff) {
if (!isnan(got)) { printf(" FAIL 0xff must be NaN, got %g\n", got); bad++; }
continue;
}
double want = ldexp(1.0, v - 127);
/* exact equality is the right test: every value is a power of two and
* therefore representable, so "close enough" would hide a real error. */
if ((double)got != want) {
if (bad < 5) printf(" FAIL v=%d: got %g, want %g\n", v, got, want);
bad++;
}
}
check(bad == 0, "all 256 UE8M0 values decode exactly");
if (bad == 0) printf(" ok all 256 UE8M0 values decode exactly\n");

/* --- the boundary that the bit trick got wrong --- */
check(ue8m0_to_f32(0) != 0.0f, "v=0 is 2^-127, NOT zero (the bit-trick bug)");
check((double)ue8m0_to_f32(0) == ldexp(1.0, -127), "v=0 == 2^-127 exactly");
printf(" ok v=0 -> %g (not zero)\n", (double)ue8m0_to_f32(0));

/* --- the values a reader is most likely to hit --- */
check(ue8m0_to_f32(127) == 1.0f, "v=127 -> 1.0");
check(ue8m0_to_f32(128) == 2.0f, "v=128 -> 2.0");
check(ue8m0_to_f32(126) == 0.5f, "v=126 -> 0.5");
/* the two scales actually observed in DeepSeek-V4's attention tensors */
check((double)ue8m0_to_f32(115) == ldexp(1.0, -12), "v=115 -> 2^-12 (real checkpoint value)");
check((double)ue8m0_to_f32(116) == ldexp(1.0, -11), "v=116 -> 2^-11 (real checkpoint value)");
printf(" ok 1.0 / 2.0 / 0.5 and the two scales seen in a real checkpoint\n");

/* --- element sizes: a wrong one here is an OOB read, not a wrong number --- */
check(st_dtype_esz(0) == 2, "BF16 is 2 bytes");
check(st_dtype_esz(1) == 2, "F16 is 2 bytes");
check(st_dtype_esz(2) == 4, "F32 is 4 bytes");
check(st_dtype_esz(3) == 1, "U8/I8 is 1 byte");
check(st_dtype_esz(4) == 1, "F8_E4M3 is 1 byte");
check(st_dtype_esz(5) == 1, "F8_E8M0 is 1 byte");
check(st_dtype_esz(6) == 8, "I64 is 8 bytes");
printf(" ok element sizes for all seven dtype codes\n");

/* --- the codes themselves must not move: containers on disk depend on the
* reader agreeing with what wrote them --- */
check(st_dtype_code("BF16") == 0 && st_dtype_code("F16") == 1 &&
st_dtype_code("F32") == 2 && st_dtype_code("U8") == 3 &&
st_dtype_code("I8") == 3, "existing dtype codes are unchanged");
check(st_dtype_code("F8_E4M3") == 4 && st_dtype_code("float8_e4m3fn") == 4,
"F8_E4M3 and its safetensors spelling both map to 4");
check(st_dtype_code("F8_E8M0") == 5, "F8_E8M0 maps to 5");
check(st_dtype_code("I64") == 6 && st_dtype_code("U64") == 6, "I64/U64 map to 6");
printf(" ok dtype codes: 0-3 unchanged, 4/5/6 added\n");

printf("test_ue8m0: %s\n", fails ? "FAILED" : "ok");
return fails ? 1 : 0;
}
Loading