Files
colibri/c/tests/test_int3.c
T
FABIOTESS 5e42e70ec3 int3-g64 (fmt=5): per-group-scale 3-bit weight format — engine, converter, tests
New weight format: int3 with ONE f32 scale per 64-input group (3.5 bits/weight
effective). Per group: 16B low plane (2 bits/val, int2 layout) + 8B high plane
(1 bit/val), values [-4,3] stored v+4. Same quantization math as
tools/quant_ablation.py _quant_last_dim(bits=3, group=64) from #132, whose
OLMoE ablation measured int3-g64 BEATING the shipped per-row int4 on quality
(-7.5 vs -9.3pp) at ~14% fewer bits.

Engine (placed per the #391 split): matmul_i3 + pack_int3_g64 + I3_* layout
helpers in quant.h next to their kernel family; fmt=5 branches in colibri.c's
qt_bytes/qt_alloc/qt_fill/matmul_qt/embed_row/qt_addrow/qt_matvec_rows.

Format detection now goes through #413's qt_resolve_fmt: fmt=5 registers its
distinct weight-byte layout O*ceil(I/64)*24 and its scale cardinality
O*ceil(I/64) there, validated against [O,I] like every other format. int3-g64
and grouped-int4-at-gs=64 carry the SAME scale count, so the weight bytes are
the int3 tag; row formats keep precedence for the small-I shapes where byte
counts coincide. The io_uring expert path still used the raw ?1:?2:3 byte
inference (it missed fmt=4 grouping entirely and never set gs) — converted to
qt_resolve_fmt like the other two expert paths.

Backends: qt_cuda_upload returns 0 for fmt=5 (tensor stays CPU-side, the
documented fallback), the dense CUDA matmul gate excludes fmt=5, and Metal's
existing fmt gates (gemm fmt<=3, moe fmt 1/2) already reject it.

Converter: quant_int3_g64 in convert_fp8_to_int4.py; --ebits 3/--xbits 3 now
emit it (previously bits=3 silently produced int4).

Tests: tests/test_int3.c (bit-exact pack/unpack vs reference, matmul_i3 vs
dequant-matmul incl. short tail groups and the real GLM I=7168, QT plumbing,
qt_resolve_fmt disambiguation incl. the same-scale-count fmt=4/fmt=5 pair,
outlier-rows RMS: int3-g64 3.3x lower error than per-row int4),
tests/test_int3_load.c (hand-rolled .safetensors fixture through st_init +
qt_from_disk: fmt=5 detected and loaded next to an int4 control tensor),
tests/test_int3_convert.py (NumPy pack round-trip vs independent decoder).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-20 17:31:05 +01:00

148 lines
6.9 KiB
C

/* int3-g64 (fmt=5) tests: pack layout, dequant round-trip vs plain-C reference,
* matmul_i3 (NEON + scalar tail) vs reference dequant-matmul, per-row helpers,
* the .qs-size format tag, and the quality claim in miniature (per-group int3
* beats per-row int4 on rows with outliers — the #132 result this format ships). */
#define main coli_glm_main_unused
#include "../colibri.c"
#undef main
#include <stdio.h>
#include <string.h>
#include <math.h>
static int fails = 0;
#define CHECK(c) do{ if(!(c)){ printf("FAIL %s:%d: %s\n", __FILE__, __LINE__, #c); fails++; } }while(0)
static uint64_t rng = 0x9E3779B97F4A7C15ull;
static float rndf(void){ rng ^= rng << 13; rng ^= rng >> 7; rng ^= rng << 17;
return ((int64_t)(rng & 0xFFFFF) - 0x80000) / (float)0x80000; }
/* reference: quantize like pack_int3_g64 but keep dequantized f32 (mirrors
* quant_ablation._quant_last_dim(bits=3, group=64)) */
static void ref_i3_dequant(const float *w, float *dq, int O, int I){
int64_t ng=i3_groups(I);
for(int o=0;o<O;o++) for(int64_t g=0;g<ng;g++){
int base=(int)(g*I3_GROUP), n=I-base<I3_GROUP?I-base:I3_GROUP;
float amax=0; for(int k=0;k<n;k++){ float a=fabsf(w[(int64_t)o*I+base+k]); if(a>amax)amax=a; }
float s=amax/3.f; if(s<1e-8f)s=1e-8f;
for(int k=0;k<n;k++){
int v=(int)lrintf(w[(int64_t)o*I+base+k]/s); if(v>3)v=3; if(v<-4)v=-4;
dq[(int64_t)o*I+base+k]=(float)v*s;
}
}
}
static void unpack_i3(const uint8_t *q3, const float *s, float *dq, int O, int I){
int64_t ng=i3_groups(I), rb=i3_rowbytes(I);
for(int o=0;o<O;o++) for(int64_t g=0;g<ng;g++){
const uint8_t *lo=q3+(int64_t)o*rb+g*I3_GBYTES, *hi=lo+16;
int base=(int)(g*I3_GROUP), n=I-base<I3_GROUP?I-base:I3_GROUP;
for(int k=0;k<n;k++){
unsigned u=((lo[k>>2]>>((k&3)*2))&3)|(((hi[k>>3]>>(k&7))&1)<<2);
dq[(int64_t)o*I+base+k]=(float)((int)u-4)*s[(int64_t)o*ng+g];
}
}
}
int main(void){
const int Is[]={64,128,192,100,65,7168}; /* incl. short tail groups and one real GLM dim */
enum { O=7, MAXI=7168 };
static float w[(int64_t)O*MAXI], dq_ref[(int64_t)O*MAXI], dq_pk[(int64_t)O*MAXI];
static float x[4*MAXI], y_ref[4*O], y_ker[4*O];
static uint8_t q3[(int64_t)O*(MAXI/64+1)*24];
static float sc[(int64_t)O*(MAXI/64+1)];
for(unsigned c=0;c<sizeof Is/sizeof *Is;c++){
int I=Is[c];
for(int64_t i=0;i<(int64_t)O*I;i++) w[i]=rndf()*0.05f;
w[3]=1.7f; w[(int64_t)2*I+5]=-2.2f; /* outliers */
/* 1. pack -> unpack == reference quantize-dequantize, bit for bit */
pack_int3_g64(w, q3, sc, O, I);
ref_i3_dequant(w, dq_ref, O, I);
unpack_i3(q3, sc, dq_pk, O, I);
int bad=0;
for(int64_t i=0;i<(int64_t)O*I;i++) if(dq_pk[i]!=dq_ref[i]) bad++;
CHECK(bad==0);
/* 2. matmul_i3 == matmul over the dequantized reference (fp tolerance:
* NEON fma order differs from the scalar reference loop) */
for(int S=1;S<=4;S+=3){
for(int64_t i=0;i<(int64_t)S*I;i++) x[i]=rndf();
matmul_i3(y_ker, x, q3, sc, S, I, O);
for(int s=0;s<S;s++) for(int o=0;o<O;o++){
double a=0; for(int i=0;i<I;i++) a+=(double)dq_ref[(int64_t)o*I+i]*x[(int64_t)s*I+i];
y_ref[s*O+o]=(float)a;
}
for(int i=0;i<S*O;i++){
float d=fabsf(y_ker[i]-y_ref[i]), m=fabsf(y_ref[i])>1?fabsf(y_ref[i]):1;
if(d/m>2e-4f){ CHECK(!"matmul_i3 mismatch"); break; }
}
}
/* 3. QT plumbing: qt_alloc(bits=3) -> qt_fill -> matmul_qt & qt_bytes & helpers */
QT t; qt_alloc(&t, O, I, 3);
CHECK(t.fmt==5);
qt_fill(&t, w, 3);
CHECK(qt_bytes(&t)==(int64_t)O*i3_rowbytes(I)+(int64_t)O*i3_groups(I)*4);
matmul_qt(y_ker, x, &t, 1);
for(int o=0;o<O;o++){
double a=0; for(int i=0;i<I;i++) a+=(double)dq_ref[(int64_t)o*I+i]*x[i];
float d=fabsf(y_ker[o]-(float)a), m=fabsf((float)a)>1?fabsf((float)a):1;
CHECK(d/m<=2e-4f);
}
float acc[MAXI]; memset(acc,0,I*sizeof(float));
qt_addrow(&t, 2, 0.5f, acc);
for(int i=0;i<I;i++) CHECK(fabsf(acc[i]-0.5f*dq_ref[(int64_t)2*I+i])<=1e-6f);
float yr[3];
qt_matvec_rows(&t, 1, 3, x, yr);
for(int j=0;j<3;j++){
double a=0; for(int i=0;i<I;i++) a+=(double)dq_ref[(int64_t)(1+j)*I+i]*x[i];
float d=fabsf(yr[j]-(float)a), m=fabsf((float)a)>1?fabsf((float)a):1;
CHECK(d/m<=2e-4f);
}
/* 4. format resolution through the #413 gate: fmt=5 is tagged by its distinct
* WEIGHT byte count. int3-g64 and grouped-int4-at-gs=64 carry the SAME scale
* cardinality O*ceil(I/64), so the pair (weight bytes, scale bytes) must
* disambiguate: same scales, int4 weights -> fmt=4/gs=64; int3 weights -> fmt=5.
* Only well-posed for I > 256: below that, O row scales legitimately match a
* 1-group grouped layout too (detect_group_size probes gs up to 256), so
* per-row vs grouped is not distinguishable from byte counts alone. */
if(I>256){ int gs=-1;
int64_t ns_g64=(int64_t)O*i3_groups(I)*4, ns_row=(int64_t)O*4;
CHECK(qt_resolve_fmt("t.i3", O, I, (int64_t)O*i3_rowbytes(I), ns_g64, &gs)==5);
CHECK(gs==0);
CHECK(qt_resolve_fmt("t.i8", O, I, (int64_t)O*I, ns_row, &gs)==1);
CHECK(qt_resolve_fmt("t.i4", O, I, (int64_t)O*((I+1)/2), ns_row, &gs)==2);
CHECK(qt_resolve_fmt("t.i4g", O, I, (int64_t)O*((I+1)/2), ns_g64, &gs)==4);
CHECK(gs==64);
CHECK(qt_resolve_fmt("t.i2", O, I, (int64_t)O*((I+3)/4), ns_row, &gs)==3); }
free(t.q4); free(t.s);
}
/* 5. quality in miniature: on rows with outliers, per-group int3 must beat
* per-row int4 on reconstruction RMS (the #132 finding this format ships). */
{
int I=1024;
for(int64_t i=0;i<(int64_t)O*I;i++) w[i]=rndf()*0.02f;
for(int o=0;o<O;o++) w[(int64_t)o*I+(o*37)%I]=1.5f; /* one outlier per row */
ref_i3_dequant(w, dq_ref, O, I);
QT t4; qt_alloc(&t4, O, I, 4); qt_fill(&t4, w, 4);
double e3=0, e4=0;
for(int o=0;o<O;o++) for(int i=0;i<I;i++){
float w4; { const uint8_t *q=t4.q4+(int64_t)o*((I+1)/2); uint8_t b=q[i>>1];
int v=(i&1)?((int)(b>>4)-8):((int)(b&0xF)-8); w4=(float)v*t4.s[o]; }
double d3=w[(int64_t)o*I+i]-dq_ref[(int64_t)o*I+i], d4=w[(int64_t)o*I+i]-w4;
e3+=d3*d3; e4+=d4*d4;
}
CHECK(e3 < e4);
printf(" outlier-rows RMS: int3-g64 %.3e < int4-row %.3e (ratio %.2f)\n",
sqrt(e3/((double)O*I)), sqrt(e4/((double)O*I)), sqrt(e4/e3));
free(t4.q4); free(t4.s);
}
if(fails){ printf("int3-g64 tests: %d FAILED\n", fails); return 1; }
printf("int3-g64 tests: ok\n");
return 0;
}