// TinyDecide device engine. See tinydecide.h and tinydecide.js (the reference). // // Memory plan ("streamed kernel"): one scratch area per call, ~2.2 KB per token. Each layer keeps // the fp32 residual x plus an int8 copy of the layer input; attention runs one head at a time with // W_o accumulated straight into x, and the FFN runs in 64-neuron chunks with fc2 accumulated into x. // Every weight row is read once per layer and applied to all tokens while it is hot in cache; rows // are split across both cores on dual-core chips. #include "tinydecide.h" #include "td_meta.h" #include "td_unicode.h" #include #include #include #include #ifdef ESP_PLATFORM #include "sdkconfig.h" #include #include #include #include #include #include #include static uint32_t now_ms() { return (uint32_t)(esp_timer_get_time() / 1000); } #if CONFIG_IDF_TARGET_ESP32S3 #define TD_PIE 1 #endif #if !CONFIG_FREERTOS_UNICORE #define TD_DUAL 1 #endif #else #include static uint32_t now_ms() { using namespace std::chrono; return (uint32_t)duration_cast(steady_clock::now().time_since_epoch()).count(); } #endif #ifndef TD_HEAP_RESERVE #define TD_HEAP_RESERVE (24 * 1024) // heap left to the rest of the app while a pass runs #endif namespace td { static const uint8_t* g_m = nullptr; // model.bin static uint32_t g_vn = 0; // vocab.bin static const uint32_t* g_voff = nullptr; static const uint16_t* g_vid = nullptr; static const char* g_vblob = nullptr; static inline const float* F(uint32_t off) { return (const float*)(g_m + off); } const char* statusText(Status s) { switch (s) { case OK: return "ok"; case NOT_READY: return "model not loaded"; case BAD_ARGS: return "bad arguments"; case QUESTION_TOO_LONG: return "question too long"; case NO_MEMORY: return "out of memory"; } return "?"; } // ============================================================================ // Tokenizer: WordPiece.encode of tinydecide.js, on UTF-8 // ============================================================================ static int vocabFind(const char* s, int len) { int lo = 0, hi = (int)g_vn - 1; while (lo <= hi) { int mid = (lo + hi) >> 1; const char* e = g_vblob + g_voff[mid]; int el = (int)(g_voff[mid + 1] - g_voff[mid]); int c = memcmp(s, e, len < el ? len : el); if (c == 0) c = len - el; if (c == 0) return g_vid[mid]; if (c < 0) hi = mid - 1; else lo = mid + 1; } return -1; } static inline bool asciiPunct(uint32_t c) { return (c >= 33 && c <= 47) || (c >= 58 && c <= 64) || (c >= 91 && c <= 96) || (c >= 123 && c <= 126); } // What the reference does to a code point >= 0x80 (td_unicode.h). For MAP, *map is set. static uint8_t classify(uint32_t cp, const tdu::Map** map) { int lo = 0, hi = tdu::N_RANGES - 1; while (lo <= hi) { int mid = (lo + hi) >> 1; if (cp < tdu::RANGES[mid].lo) hi = mid - 1; else if (cp > tdu::RANGES[mid].hi) lo = mid + 1; else return tdu::RANGES[mid].kind; } lo = 0; hi = tdu::N_MAPS - 1; while (lo <= hi) { int mid = (lo + hi) >> 1; if (cp < tdu::MAPS[mid].cp) hi = mid - 1; else if (cp > tdu::MAPS[mid].cp) lo = mid + 1; else { *map = &tdu::MAPS[mid]; return tdu::MAP; } } return tdu::WORD; } // Decodes one UTF-8 code point from at most rem bytes; invalid input becomes U+FFFD (dropped, like // the reference). static uint32_t nextCp(const unsigned char* p, size_t rem, int* len) { const unsigned char c = p[0]; if (c < 0x80) { *len = 1; return c; } const int n = c >= 0xC2 && c <= 0xDF ? 2 : c >= 0xE0 && c <= 0xEF ? 3 : c >= 0xF0 && c <= 0xF4 ? 4 : 0; if (n == 0 || (size_t)n > rem) { *len = 1; return 0xFFFD; } uint32_t cp = c & (0x7F >> n); for (int i = 1; i < n; i++) { if ((p[i] & 0xC0) != 0x80) { *len = 1; return 0xFFFD; } cp = (cp << 6) | (p[i] & 0x3F); } *len = n; if ((n == 3 && cp < 0x800) || (n == 4 && (cp < 0x10000 || cp > 0x10FFFF)) || (cp >= 0xD800 && cp <= 0xDFFF)) return 0xFFFD; return cp; } static int putUtf8(uint32_t cp, char* o) { if (cp < 0x80) { o[0] = (char)cp; return 1; } if (cp < 0x800) { o[0] = (char)(0xC0 | (cp >> 6)); o[1] = (char)(0x80 | (cp & 0x3F)); return 2; } if (cp < 0x10000) { o[0] = (char)(0xE0 | (cp >> 12)); o[1] = (char)(0x80 | ((cp >> 6) & 0x3F)); o[2] = (char)(0x80 | (cp & 0x3F)); return 3; } o[0] = (char)(0xF0 | (cp >> 18)); o[1] = (char)(0x80 | ((cp >> 12) & 0x3F)); o[2] = (char)(0x80 | ((cp >> 6) & 0x3F)); o[3] = (char)(0x80 | (cp & 0x3F)); return 4; } namespace { struct Tok { uint16_t* ids; int32_t* starts; int32_t* ends; int max, n; // current word: normalised code points with their source byte ranges uint32_t cp[tdm::MAX_CHARS]; int32_t a[tdm::MAX_CHARS], b[tdm::MAX_CHARS]; int wl; bool longWord; int32_t firstA, lastB; char utf[tdm::MAX_CHARS * 4]; // the word as UTF-8 char piece[tdm::MAX_CHARS * 4 + 2]; uint16_t uo[tdm::MAX_CHARS + 1]; // byte offset of each code point in utf uint16_t pid[tdm::MAX_CHARS]; int32_t pa[tdm::MAX_CHARS], pb[tdm::MAX_CHARS]; void emit(int id, int32_t s, int32_t e) { if (n >= max) return; ids[n] = (uint16_t)id; if (starts) starts[n] = s; if (ends) ends[n] = e; n++; } void add(uint32_t c, int32_t s, int32_t e) { if (wl == 0 && !longWord) firstA = s; lastB = e; if (wl < tdm::MAX_CHARS) { cp[wl] = c; a[wl] = s; b[wl] = e; wl++; } else longWord = true; } void flush() { if (longWord) emit(tdm::SP_UNK, firstA, lastB); else if (wl) encodeWord(); wl = 0; longWord = false; } void isolated(uint32_t c, int32_t s, int32_t e) { flush(); add(c, s, e); flush(); } void encodeWord() { int ub = 0; for (int i = 0; i < wl; i++) { uo[i] = (uint16_t)ub; ub += putUtf8(cp[i], utf + ub); } uo[wl] = (uint16_t)ub; int np = 0, start = 0; while (start < wl) { int end = wl, got = -1; while (start < end) { if (start == 0) got = vocabFind(utf, uo[end]); else { // continuation piece: "##" + the code points const int len = uo[end] - uo[start]; piece[0] = '#'; piece[1] = '#'; memcpy(piece + 2, utf + uo[start], len); got = vocabFind(piece, len + 2); } if (got >= 0) break; end--; } if (got < 0) { emit(tdm::SP_UNK, a[0], b[wl - 1]); return; } pid[np] = (uint16_t)got; pa[np] = a[start]; pb[np] = b[end - 1]; np++; start = end; } for (int i = 0; i < np; i++) emit(pid[i], pa[i], pb[i]); } }; Tok g_tok; } // namespace int tokenize(const char* text, uint16_t* ids, int max, int32_t* starts, int32_t* ends) { return tokenize(text, text ? strlen(text) : 0, ids, max, starts, ends); } int tokenize(const char* text, size_t n, uint16_t* ids, int max, int32_t* starts, int32_t* ends) { if (!g_m || !text || max <= 0) return 0; Tok& t = g_tok; t.ids = ids; t.starts = starts; t.ends = ends; t.max = max; t.n = 0; t.wl = 0; t.longWord = false; const unsigned char* p = (const unsigned char*)text; int32_t i = 0; while ((size_t)i < n && t.n < max) { int len; const uint32_t c = nextCp(p + i, n - (size_t)i, &len); const int32_t s = i, e = i + len; i = e; if (c < 0x80) { if (c == ' ' || c == '\t' || c == '\n' || c == '\r') { t.flush(); continue; } if (c < 32 || c == 127) continue; // control characters are dropped if (asciiPunct(c)) { t.isolated(c, s, e); continue; } t.add(c >= 'A' && c <= 'Z' ? c + 32 : c, s, e); continue; } if (c >= 0xAC00 && c <= 0xD7A3) { // Hangul syllable: NFD to 2-3 jamo const uint32_t k = c - 0xAC00, tj = k % 28; t.add(0x1100 + k / 588, s, e); t.add(0x1161 + (k % 588) / 28, s, e); if (tj) t.add(0x11A7 + tj, s, e); continue; } const tdu::Map* m = nullptr; switch (classify(c, &m)) { case tdu::DROP: case tdu::MN: break; case tdu::SPACE: t.flush(); break; case tdu::CJK: case tdu::PUNCT: t.isolated(c, s, e); break; case tdu::MAP: for (int j = 0; j < m->len; j++) { const uint32_t o = tdu::MAP_OUT[m->off + j]; if (o & tdu::OUT_PUNCT) t.isolated(o & ~tdu::OUT_PUNCT, s, e); else t.add(o, s, e); } break; default: t.add(c, s, e); } } t.flush(); return t.n < max ? t.n : max; } // ============================================================================ // Kernels // ============================================================================ static inline float bf16(const uint8_t* p) { uint32_t u = (uint32_t)(p[0] | (p[1] << 8)) << 16; float f; memcpy(&f, &u, 4); return f; } static inline float dotq(const uint8_t* nib, const int8_t* xq, const uint8_t* sc, const float* xs, int nb) { #if defined(TD_PIE) return dot_q4q8_pie(nib, xq, sc, xs, nb); #else float acc = 0.0f; for (int b = 0; b < nb; b++) { int isum = 0; for (int k = 0; k < 16; k++) { uint8_t bk = nib[k]; isum += ((int)(bk & 0x0F) - 8) * (int)xq[k]; isum += ((int)(bk >> 4) - 8) * (int)xq[k + 16]; } acc += (float)isum * (bf16(sc) * xs[b]); nib += 16; xq += 32; sc += 2; } return acc; #endif } // Symmetric int8 per 32-element block, one fp32 scale each (llama.cpp Q8_0). static void quant8(const float* x, int8_t* xq, float* xs, int n) { for (int b = 0; b < n / 32; b++) { float mx = 0.0f; for (int i = 0; i < 32; i++) { float a = fabsf(x[i]); if (a > mx) mx = a; } float s = mx / 127.0f, inv = s > 0 ? 1.0f / s : 0.0f; xs[b] = s; for (int i = 0; i < 32; i++) xq[i] = (int8_t)lrintf(x[i] * inv); x += 32; xq += 32; } } // out[t*os + (r - r0)] (=|+=) W[r, blk0*32 .. (blk0+nbk)*32) . x[t] + bias[r], for r in [r0, r1). struct MM { const tdm::Q4* w; int blk0, nbk; const int8_t* xq; int xqs; // int8 activations, stride per token (bytes) const float* xs; int xss; // block scales, stride per token int T; float* out; int os; const float* bias; bool acc; int r0; }; static void mmRows(const MM& a, int ra, int rb) { const int nbRow = a.w->cols / 32; for (int r = ra; r < rb; r++) { const uint8_t* nib = g_m + a.w->nib + ((size_t)r * nbRow + a.blk0) * 16; const uint8_t* sc = g_m + a.w->sc + ((size_t)r * nbRow + a.blk0) * 2; const float b = a.bias ? a.bias[r] : 0.0f; float* o = a.out + (r - a.r0); for (int t = 0; t < a.T; t++) { float v = dotq(nib, a.xq + (size_t)t * a.xqs, sc, a.xs + (size_t)t * a.xss, a.nbk) + b; if (a.acc) o[(size_t)t * a.os] += v; else o[(size_t)t * a.os] = v; } } } // parallel(fn, ctx, n): fn(ctx, a, b) over [0, n), split in two halves across both cores. typedef void (*RangeFn)(void* ctx, int a, int b); #if defined(TD_DUAL) // Persistent worker on the other core; per call it costs two semaphore operations. static SemaphoreHandle_t s_go = nullptr, s_done = nullptr; static RangeFn s_fn = nullptr; static void* s_ctx = nullptr; static int s_ja = 0, s_jb = 0; static int s_core = -1; static void worker(void*) { for (;;) { xSemaphoreTake(s_go, portMAX_DELAY); s_fn(s_ctx, s_ja, s_jb); xSemaphoreGive(s_done); } } static void parallel(RangeFn fn, void* ctx, int n) { if (!s_go) { s_go = xSemaphoreCreateBinary(); s_done = xSemaphoreCreateBinary(); s_core = xPortGetCoreID() == 0 ? 1 : 0; xTaskCreatePinnedToCore(worker, "td_mm", 3072, nullptr, 3, nullptr, s_core); } if (xPortGetCoreID() == s_core || n < 2) { fn(ctx, 0, n); return; } // caller moved cores: run alone const int split = n / 2; s_fn = fn; s_ctx = ctx; s_ja = split; s_jb = n; xSemaphoreGive(s_go); fn(ctx, 0, split); xSemaphoreTake(s_done, portMAX_DELAY); } #else static void parallel(RangeFn fn, void* ctx, int n) { fn(ctx, 0, n); } #endif #ifdef ESP_PLATFORM static void* scratchAlloc(size_t n) { void* p = heap_caps_aligned_alloc(16, n, MALLOC_CAP_8BIT | MALLOC_CAP_INTERNAL); #if CONFIG_SPIRAM if (!p) p = heap_caps_aligned_alloc(16, n, MALLOC_CAP_8BIT | MALLOC_CAP_SPIRAM); #endif return p; } static void scratchFree(void* p) { heap_caps_free(p); } static bool reserveOk() { return heap_caps_get_free_size(MALLOC_CAP_8BIT) >= TD_HEAP_RESERVE; } #else static void* scratchAlloc(size_t n) { return aligned_alloc(16, (n + 15) & ~(size_t)15); } static void scratchFree(void* p) { free(p); } static bool reserveOk() { return true; } #endif static void mmRange(void* ctx, int a, int b) { const MM& m = *(const MM*)ctx; mmRows(m, m.r0 + a, m.r0 + b); } static void mm(const MM& a, int r1) { parallel(mmRange, (void*)&a, r1 - a.r0); } static void layernorm(float* x, int d, const float* w, const float* b, float eps) { float m = 0; for (int i = 0; i < d; i++) m += x[i]; m /= d; float v = 0; for (int i = 0; i < d; i++) { float z = x[i] - m; v += z * z; } v /= d; const float r = 1.0f / sqrtf(v + eps); for (int i = 0; i < d; i++) x[i] = (x[i] - m) * r * w[i] + b[i]; } // Exact (erf) GELU from a table with linear interpolation; step 1/64 over [-8, 8], error < 3e-5. static constexpr int GELU_N = 1024; static float g_gelu[GELU_N + 1]; static void geluInit() { for (int i = 0; i <= GELU_N; i++) { const double x = -8.0 + 16.0 * i / GELU_N; g_gelu[i] = (float)(0.5 * x * (1.0 + erf(x * 0.70710678118654752))); } } static inline float gelu(float x) { if (x >= 8.0f) return x; if (x <= -8.0f) return 0.0f; const float f = (x + 8.0f) * (GELU_N / 16.0f); const int i = (int)f; const float w = f - (float)i; return g_gelu[i] + (g_gelu[i + 1] - g_gelu[i]) * w; } // y = W v for an int8 head matrix (one f32 scale per row). static void mvI8(const tdm::I8& w, const float* v, float* y) { const int8_t* q = (const int8_t*)(g_m + w.q); const float* s = F(w.sc); for (int r = 0; r < w.rows; r++) { float acc = 0; const int8_t* row = q + (size_t)r * w.cols; for (int c = 0; c < w.cols; c++) acc += (float)row[c] * v[c]; y[r] = acc * s[r]; } } // ============================================================================ // Request layout (encodeRequest in tinydecide.js) // ============================================================================ struct Req { uint16_t* ids; // [T] uint16_t* pos; // [T] uint8_t* blk; // [T] 0 = state, k = question k int T, S; // total tokens; state tokens incl. the <|state|> marker int nq; int qs[MAX_QUESTIONS + 1], qe[MAX_QUESTIONS + 1]; // block k occupies [qs[k], qe[k]) int ans[MAX_QUESTIONS]; int opt[MAX_QUESTIONS][MAX_OPTIONS]; int32_t sa[tdm::TS_MAX], sb[tdm::TS_MAX]; // byte range of each state text token bool truncated; }; // Question block: <|type|> text [<|sep|> (option <|o|> or <|lv|>)*] <|ans|>. Writes up to cap ids // and returns the block's full length, which may exceed cap (then the question is too long). static int questionBlock(const Question& q, uint16_t* ids, int cap, int* optLocal, int* ansLocal) { static uint16_t tmp[tdm::Q_MAX + 1]; int m = 0; auto put = [&](int id) { if (m < cap) ids[m] = (uint16_t)id; m++; }; put(tdm::SP_TYPE[q.type]); int k = tokenize(q.text, tmp, tdm::Q_MAX + 1); for (int i = 0; i < k; i++) put(tmp[i]); if (q.n_options > 0) { put(tdm::SP_SEP); const int mk = q.type == CHOICE ? tdm::SP_O : tdm::SP_LV; for (int i = 0; i < q.n_options; i++) { k = tokenize(q.options[i] ? q.options[i] : "", tmp, tdm::Q_MAX + 1); for (int j = 0; j < k; j++) put(tmp[j]); if (optLocal) optLocal[i] = m; put(mk); if (m > tdm::Q_MAX) return m; } } if (ansLocal) *ansLocal = m; put(tdm::SP_ANS); return m; } static Status checkArgs(const Question* qs, int nq) { if (nq < 1 || nq > MAX_QUESTIONS || !qs) return BAD_ARGS; for (int k = 0; k < nq; k++) { const Question& q = qs[k]; if (q.type > SPAN || !q.text) return BAD_ARGS; if ((q.type == CHOICE || q.type == SCORE) && (q.n_options < 2 || q.n_options > MAX_OPTIONS || !q.options)) return BAD_ARGS; if (q.n_options < 0 || q.n_options > MAX_OPTIONS || (q.n_options && !q.options)) return BAD_ARGS; if (q.protos && (q.type == SPAN || !q.protos->vec || !q.protos->cnt)) return BAD_ARGS; } return OK; } int requestTokens(const Question* qs, int nq, int state_tokens) { if (checkArgs(qs, nq) != OK) return -1; static uint16_t ids[tdm::Q_MAX]; int T = 1 + (state_tokens < STATE_MAX ? state_tokens : STATE_MAX); for (int k = 0; k < nq; k++) { const int n = questionBlock(qs[k], ids, tdm::Q_MAX, nullptr, nullptr); if (n > tdm::Q_MAX) return -1; T += n; } return T; } // ============================================================================ // Encoder // ============================================================================ // Several blocks rather than one, so the scratch fits a fragmented heap (no PSRAM on most boards): // A: x [T, D] f32 B: Qh, Kh, Vh [T, 64] f32 C: xq, hq int8 + block scales + scores. struct Scratch { uint8_t* a; uint8_t* q; uint8_t* k; uint8_t* v; uint8_t* c; }; static size_t bytesA(int T) { return (size_t)T * tdm::D * sizeof(float); } static size_t bytesB(int T) { return (size_t)T * tdm::DHEAD * sizeof(float); } // each of Q, K, V static size_t bytesC(int T) { return (size_t)T * (tdm::D + tdm::DHEAD + (tdm::D / 32 + tdm::DHEAD / 32 + 2) * sizeof(float)); } size_t scratchBytes(int T) { return bytesA(T) + 3 * bytesB(T) + bytesC(T); } static void freeScratch(Scratch& m) { for (uint8_t* p : {m.c, m.v, m.k, m.q, m.a}) if (p) scratchFree(p); m = {}; } static bool allocScratch(Scratch& m, int T) { m = {}; m.a = (uint8_t*)scratchAlloc(bytesA(T)); // largest first m.q = m.a ? (uint8_t*)scratchAlloc(bytesB(T)) : nullptr; m.k = m.q ? (uint8_t*)scratchAlloc(bytesB(T)) : nullptr; m.v = m.k ? (uint8_t*)scratchAlloc(bytesB(T)) : nullptr; m.c = m.v ? (uint8_t*)scratchAlloc(bytesC(T)) : nullptr; if (m.c && reserveOk()) return true; freeScratch(m); return false; } // GELU over one FFN chunk for tokens [a, b), then int8 for the fc2 slice. struct GeluCtx { float* Hc; int8_t* hq; float* hs; }; static void geluRange(void* ctx, int a, int b) { const GeluCtx& g = *(const GeluCtx*)ctx; const int DHd = tdm::DHEAD; for (int t = a; t < b; t++) { float* hrow = g.Hc + (size_t)t * DHd; for (int c = 0; c < DHd; c++) hrow[c] = gelu(hrow[c]); quant8(hrow, g.hq + (size_t)t * DHd, g.hs + (size_t)t * (DHd / 32), DHd); } } // One head's attention for query rows [a, b): softmax(q k / sqrt(64)) v, written as int8 to hq. // State tokens see the state; a question's tokens see the state and their own block. struct AttnCtx { const float *Qh, *Kh, *Vh; int8_t* hq; float* hs; float* sc; // [2, T] scores, one row per core const Req* r; float scale; }; static void attnRange(void* ctx, int a, int b) { const AttnCtx& c = *(const AttnCtx*)ctx; const Req& r = *c.r; const int DHd = tdm::DHEAD; float* sc = c.sc + (a == 0 ? 0 : r.T); float out[64]; for (int i = a; i < b; i++) { const int k = r.blk[i]; const int ranges[2][2] = {{0, r.S}, {k ? r.qs[k] : 0, k ? r.qe[k] : 0}}; const float* q = c.Qh + (size_t)i * DHd; float mx = -1e30f; int n = 0; for (const auto& rg : ranges) for (int j = rg[0]; j < rg[1]; j++) { const float* kk = c.Kh + (size_t)j * DHd; float s = 0; for (int d = 0; d < DHd; d++) s += q[d] * kk[d]; s *= c.scale; sc[n++] = s; if (s > mx) mx = s; } float z = 0; for (int j = 0; j < n; j++) { sc[j] = expf(sc[j] - mx); z += sc[j]; } const float iz = 1.0f / z; for (int d = 0; d < DHd; d++) out[d] = 0; n = 0; for (const auto& rg : ranges) for (int j = rg[0]; j < rg[1]; j++) { const float p = sc[n++] * iz; const float* v = c.Vh + (size_t)j * DHd; for (int d = 0; d < DHd; d++) out[d] += p * v[d]; } quant8(out, c.hq + (size_t)i * DHd, c.hs + (size_t)i * (DHd / 32), DHd); } } static void encode(const Req& r, const Scratch& m) { const int T = r.T, D = tdm::D, DHd = tdm::DHEAD, E = tdm::EMB; float* x = (float*)m.a; // [T, D] float* Qh = (float*)m.q; // [T, 64] each float* Kh = (float*)m.k; float* Vh = (float*)m.v; int8_t* xq = (int8_t*)m.c; // [T, D] (16-aligned: offsets are multiples of 64) int8_t* hq = xq + (size_t)T * D; // [T, 64] float* xs = (float*)(hq + (size_t)T * DHd); // [T, D/32] float* hs = xs + (size_t)T * (D / 32); // [T, 2] float* sc = hs + (size_t)T * (DHd / 32); // [2, T] // ---- embeddings: word (Q4) + position (int8) + type0, LayerNorm, project to D const int8_t* P = (const int8_t*)(g_m + tdm::POS.q); const float* Ps = F(tdm::POS.sc); const float* t0 = F(tdm::TYPE0); for (int t = 0; t < T; t++) { float o[tdm::EMB]; const int nb = E / 32; const uint8_t* nib = g_m + tdm::WORD.nib + (size_t)r.ids[t] * nb * 16; const uint8_t* ws = g_m + tdm::WORD.sc + (size_t)r.ids[t] * nb * 2; for (int b = 0; b < nb; b++) { const float d = bf16(ws + 2 * b); for (int k = 0; k < 16; k++) { const uint8_t byte = nib[b * 16 + k]; o[b * 32 + k] = ((int)(byte & 15) - 8) * d; o[b * 32 + k + 16] = ((int)(byte >> 4) - 8) * d; } } const int8_t* pr = P + (size_t)r.pos[t] * E; const float ps = Ps[r.pos[t]]; for (int c = 0; c < E; c++) o[c] += pr[c] * ps + t0[c]; layernorm(o, E, F(tdm::ELN_W), F(tdm::ELN_B), tdm::LN_EPS); quant8(o, xq + (size_t)t * E, xs + (size_t)t * (E / 32), E); } mm({&tdm::PROJ, 0, E / 32, xq, E, xs, E / 32, T, x, D, F(tdm::PROJ_B), false, 0}, D); const float scale = 1.0f / sqrtf((float)DHd); for (int l = 0; l < tdm::LAYERS; l++) { const tdm::Block& B = tdm::BLOCKS[l]; for (int t = 0; t < T; t++) quant8(x + (size_t)t * D, xq + (size_t)t * D, xs + (size_t)t * (D / 32), D); // ---- attention, one head at a time; W_o slice accumulated into the residual for (int h = 0; h < tdm::HEADS; h++) { const int r0 = h * DHd, r1 = r0 + DHd; mm({&B.q, 0, D / 32, xq, D, xs, D / 32, T, Qh, DHd, F(B.qb), false, r0}, r1); mm({&B.k, 0, D / 32, xq, D, xs, D / 32, T, Kh, DHd, F(B.kb), false, r0}, r1); mm({&B.v, 0, D / 32, xq, D, xs, D / 32, T, Vh, DHd, F(B.vb), false, r0}, r1); AttnCtx ac{Qh, Kh, Vh, hq, hs, sc, &r, scale}; parallel(attnRange, &ac, T); mm({&B.o, r0 / 32, DHd / 32, hq, DHd, hs, DHd / 32, T, x, D, h == 0 ? F(B.ob) : nullptr, true, 0}, D); } for (int t = 0; t < T; t++) layernorm(x + (size_t)t * D, D, F(B.ln1w), F(B.ln1b), tdm::LN_EPS); // ---- FFN in 64-neuron chunks; fc2 slice accumulated into the residual for (int t = 0; t < T; t++) quant8(x + (size_t)t * D, xq + (size_t)t * D, xs + (size_t)t * (D / 32), D); float* Hc = Qh; for (int c0 = 0; c0 < tdm::FFN; c0 += DHd) { mm({&B.fc, 0, D / 32, xq, D, xs, D / 32, T, Hc, DHd, F(B.fcb), false, c0}, c0 + DHd); GeluCtx gc{Hc, hq, hs}; parallel(geluRange, &gc, T); mm({&B.fc2, c0 / 32, DHd / 32, hq, DHd, hs, DHd / 32, T, x, D, c0 == 0 ? F(B.fc2b) : nullptr, true, 0}, D); } for (int t = 0; t < T; t++) layernorm(x + (size_t)t * D, D, F(B.ln2w), F(B.ln2b), tdm::LN_EPS); } } // ============================================================================ // Heads (TinyDecide.answer in tinydecide.js) // ============================================================================ static int bucketK(int c) { int b = 0; for (int e : {1, 2, 4, 8}) if (c > e) b++; return b; } // row t of x through the head LayerNorm static void headRow(const float* x, int t, float* h) { memcpy(h, x + (size_t)t * tdm::D, tdm::D * sizeof(float)); layernorm(h, tdm::D, F(tdm::H_NORM_W), F(tdm::H_NORM_B), 1e-5f); } static float dotf(const float* a, const float* b, int n) { float s = 0; for (int i = 0; i < n; i++) s += a[i] * b[i]; return s; } // The correction term of option i: lam * beta(cnt) * cos(qvec - c, vec_i - c). Without a center // (the older recipe) it is beta(cnt) * cos(legacy, vec_i), where legacy is the question's own // projection (h.A for choice / score, h.noul_q for noul), and lam is not applied. static float protoTerm(const Protos& p, const float* qvec, const float* legacy, int i) { if (p.cnt[i] <= 0) return 0.0f; const float* pv = p.vec + (size_t)i * QDIM; if (!p.center) { float dd = 0, na = 0, nb = 0; for (int j = 0; j < QDIM; j++) { dd += legacy[j] * pv[j]; na += legacy[j] * legacy[j]; nb += pv[j] * pv[j]; } const float sa = sqrtf(na) > 0 ? sqrtf(na) : 1e-12f, sb = sqrtf(nb) > 0 ? sqrtf(nb) : 1e-12f; return tdm::BETA[bucketK(p.cnt[i])] * dd / (sa * sb); } const float* c = p.center; float dd = 0, na = 0, nb = 0; for (int j = 0; j < QDIM; j++) { const float a = qvec[j] - c[j], b = pv[j] - c[j]; dd += a * b; na += a * a; nb += b * b; } const float sa = sqrtf(na) > 0 ? sqrtf(na) : 1e-12f, sb = sqrtf(nb) > 0 ? sqrtf(nb) : 1e-12f; return p.lam * tdm::BETA[bucketK(p.cnt[i])] * dd / (sa * sb); } static void headChoice(const Req& r, const float* x, int k, const Question& q, Answer& out) { const int sel = q.type == SCORE ? 1 : 0, n = q.n_options, DH = tdm::DH; const float Tt = tdm::TEMP[sel ? 2 : 0], s = expf(F(tdm::H_SCALE)[sel]); float h[tdm::D], qa[tdm::DH], ov[tdm::DH], L[MAX_OPTIONS]; headRow(x, r.ans[k], h); mvI8(tdm::H_A[sel], h, qa); mvI8(tdm::H_P, h, out.qvec); for (int i = 0; i < n; i++) { headRow(x, r.opt[k][i], h); mvI8(tdm::H_O[sel], h, ov); out.z0[i] = s * dotf(qa, ov, DH) / sqrtf((float)DH) / Tt; L[i] = out.z0[i] + (q.bias ? q.bias[i] : 0.0f); if (q.protos) L[i] += protoTerm(*q.protos, out.qvec, qa, i) / Tt; } float mx = L[0]; for (int i = 1; i < n; i++) if (L[i] > mx) mx = L[i]; float Z = 0; for (int i = 0; i < n; i++) { out.probs[i] = expf(L[i] - mx); Z += out.probs[i]; } float H = 0, sc = 0; out.pick = 0; for (int i = 0; i < n; i++) { out.probs[i] /= Z; if (out.probs[i] > out.probs[out.pick]) out.pick = i; if (out.probs[i] > 0) H -= out.probs[i] * logf(out.probs[i]); sc += out.probs[i] * i / (float)(n - 1); } out.n = n; out.confidence = 1.0f - H / logf((float)n); out.score = q.type == SCORE ? sc : 0.0f; } static void headNoul(const Req& r, const float* x, int k, const Question& q, Answer& out) { const float Tt = tdm::TEMP[1]; float h[tdm::D]; headRow(x, r.ans[k], h); mvI8(tdm::H_P, h, out.qvec); const float z = dotf(F(tdm::H_NOUL_W), h, tdm::D) + F(tdm::H_NOUL_B)[0]; out.z0[0] = 0.0f; out.z0[1] = z / Tt; float L = out.z0[1] + (q.bias ? q.bias[0] : 0.0f); if (q.protos) { float nq[QDIM]; if (!q.protos->center) mvI8(tdm::H_NOUL_Q, h, nq); L += (protoTerm(*q.protos, out.qvec, nq, 1) - protoTerm(*q.protos, out.qvec, nq, 0)) / Tt; } out.n = 2; out.p = 1.0f / (1.0f + expf(-L)); out.probs[0] = 1.0f - out.p; out.probs[1] = out.p; out.pick = out.p > 0.5f ? 1 : 0; } static void headSpan(const Req& r, const float* x, int k, const char* state, Answer& out) { const int S = r.S - 1, DH = tdm::DH; // state text tokens at rows 1..S const float Tt = tdm::TEMP[3], isq = 1.0f / sqrtf((float)DH); float h[tdm::D], hAns[tdm::D], qS[tdm::DH], kk[tdm::DH], qe[tdm::DH]; headRow(x, r.ans[k], hAns); mvI8(tdm::H_P, hAns, out.qvec); mvI8(tdm::H_SQ, hAns, qS); // start logits: null, then each state token; log-softmax at temperature Tt const float zNull = dotf(F(tdm::H_SNULL_W), hAns, tdm::D) + F(tdm::H_SNULL_B)[0]; static float zs[tdm::TS_MAX]; float mx = zNull; for (int t = 0; t < S; t++) { headRow(x, 1 + t, h); mvI8(tdm::H_SK, h, kk); zs[t] = dotf(qS, kk, DH) * isq; if (zs[t] > mx) mx = zs[t]; } float Z = expf((zNull - mx) / Tt); for (int t = 0; t < S; t++) Z += expf((zs[t] - mx) / Tt); const float lZ = logf(Z); int best = 0; for (int t = 1; t < S; t++) if (zs[t] > zs[best]) best = t; out.p_present = 1.0f - expf((zNull - mx) / Tt - lZ); out.tok[0] = out.tok[1] = best; out.start = out.end = 0; out.p_span = 0.0f; out.n = 0; if (S == 0) return; const float lsBest = (zs[best] - mx) / Tt - lZ; // end logits inside [best, best + SPAN_MAX) float es[tdm::DH]; mvI8(tdm::H_EQ, hAns, qe); headRow(x, 1 + best, h); mvI8(tdm::H_ES, h, es); for (int i = 0; i < DH; i++) qe[i] += es[i]; const int e1 = best + tdm::SPAN_MAX < S ? best + tdm::SPAN_MAX : S; float en[tdm::SPAN_MAX], emx = -1e30f; for (int t = best; t < e1; t++) { headRow(x, 1 + t, h); mvI8(tdm::H_EK, h, kk); en[t - best] = dotf(qe, kk, DH) * isq; if (en[t - best] > emx) emx = en[t - best]; } float EZ = 0; for (int t = best; t < e1; t++) EZ += expf((en[t - best] - emx) / Tt); int bestE = best; for (int t = best; t < e1; t++) if (en[t - best] > en[bestE - best]) bestE = t; out.tok[1] = bestE; out.p_span = expf(lsBest + (en[bestE - best] - emx) / Tt - logf(EZ)); int a = r.sa[best], b = r.sb[bestE]; auto ws = [](char c) { return c == ' ' || c == '\t' || c == '\n' || c == '\r'; }; while (a < b && ws(state[a])) a++; while (b > a && ws(state[b - 1])) b--; out.start = a; out.end = b; } // ============================================================================ // Public API // ============================================================================ bool init(const uint8_t* model, size_t model_len, const uint8_t* vocab, size_t vocab_len) { g_m = nullptr; if (!model || model_len < tdm::MODEL_BYTES || ((uintptr_t)model & 15)) return false; uint32_t fnv = 0x811c9dc5u; for (size_t i = 0; i < 65536 && i < tdm::MODEL_BYTES; i++) { fnv ^= model[i]; fnv *= 0x01000193u; } if (fnv != tdm::MODEL_FNV64K) return false; // a different model.bin than td_meta.h if (!vocab || vocab_len < 12 || memcmp(vocab, "TDV1", 4) != 0) return false; uint32_t n, blob; memcpy(&n, vocab + 4, 4); memcpy(&blob, vocab + 8, 4); if (12 + 4 * (size_t)(n + 1) + 2 * (size_t)n + blob > vocab_len) return false; geluInit(); g_vn = n; g_voff = (const uint32_t*)(vocab + 12); g_vid = (const uint16_t*)(vocab + 12 + 4 * (n + 1)); g_vblob = (const char*)(vocab + 12 + 4 * (n + 1) + 2 * n); g_m = model; return true; } #ifdef ESP_PLATFORM #if defined(TINYDECIDE_EMBED) extern "C" const uint8_t td_model_bin[], td_model_bin_end[], td_vocab_bin[], td_vocab_bin_end[]; bool initEmbedded() { return init(td_model_bin, (size_t)(td_model_bin_end - td_model_bin), td_vocab_bin, (size_t)(td_vocab_bin_end - td_vocab_bin)); } #else bool initEmbedded() { return false; } #endif bool initPartition(const char* label) { const esp_partition_t* part = esp_partition_find_first(ESP_PARTITION_TYPE_DATA, ESP_PARTITION_SUBTYPE_ANY, label); if (!part) { ESP_LOGE("tinydecide", "no data partition named '%s'", label); return false; } const void* ptr = nullptr; esp_partition_mmap_handle_t handle; if (esp_partition_mmap(part, 0, part->size, ESP_PARTITION_MMAP_DATA, &ptr, &handle) != ESP_OK) { ESP_LOGE("tinydecide", "could not map partition '%s' (%u bytes)", label, (unsigned)part->size); return false; } const uint8_t* base = (const uint8_t*)ptr; const size_t vo = (tdm::MODEL_BYTES + 15) & ~(size_t)15; if (part->size < vo + 12 || !init(base, tdm::MODEL_BYTES, base + vo, part->size - vo)) { ESP_LOGE("tinydecide", "partition '%s' does not hold this build's tinydecide-esp32.bin", label); esp_partition_munmap(handle); return false; } return true; } #endif #ifdef TD_TEST // Host tests only: the token ids of the last answer() call. static uint16_t g_lastIds[tdm::TS_MAX + MAX_QUESTIONS * tdm::Q_MAX]; static int g_lastT = 0; int debugLastIds(const uint16_t** ids) { *ids = g_lastIds; return g_lastT; } #endif Status answer(const char* state, const Question* qs, int nq, Answer* out, Info* info, int state_max) { return answer(state, state ? strlen(state) : 0, qs, nq, out, info, state_max); } Status answer(const char* state, size_t state_len, const Question* qs, int nq, Answer* out, Info* info, int state_max) { if (!g_m) return NOT_READY; if (!state || !out) return BAD_ARGS; const Status bad = checkArgs(qs, nq); if (bad != OK) return bad; const uint32_t t0 = now_ms(); if (state_max < 1) state_max = 1; if (state_max > STATE_MAX) state_max = STATE_MAX; Req* r = (Req*)malloc(sizeof(Req)); if (!r) return NO_MEMORY; // state tokens (the reference reads up to ts_max - 1 = 127) static uint16_t st[tdm::TS_MAX]; static int32_t sa[tdm::TS_MAX], sb[tdm::TS_MAX]; const int ns = tokenize(state, state_len, st, tdm::TS_MAX, sa, sb); // question blocks: lengths first, so the sequence buffers can be sized exactly int qlen = 0; static uint16_t qids[tdm::Q_MAX]; for (int k = 0; k < nq; k++) { const int n = questionBlock(qs[k], qids, tdm::Q_MAX, nullptr, nullptr); if (n > tdm::Q_MAX) { free(r); return QUESTION_TOO_LONG; } qlen += n; } const int Tmax = 1 + tdm::TS_MAX + qlen; uint8_t* seq = (uint8_t*)malloc((size_t)Tmax * 5); if (!seq) { free(r); return NO_MEMORY; } r->ids = (uint16_t*)seq; r->pos = r->ids + Tmax; r->blk = (uint8_t*)(r->pos + Tmax); r->nq = nq; // When the heap is too fragmented for the whole state, keep fewer of its tokens and try again. Scratch mem = {}; for (int keepMax = state_max;; keepMax -= 8) { if (keepMax < 1) keepMax = 1; const int keep = ns < keepMax ? ns : keepMax; r->truncated = ns > keep; r->ids[0] = tdm::SP_STATE; r->pos[0] = 0; r->blk[0] = 0; for (int i = 0; i < keep; i++) { r->ids[1 + i] = st[i]; r->pos[1 + i] = (uint16_t)(1 + i); r->blk[1 + i] = 0; r->sa[i] = sa[i]; r->sb[i] = sb[i]; } r->S = 1 + keep; int T = r->S; for (int k = 0; k < nq; k++) { int ansLocal = 0; const int n = questionBlock(qs[k], r->ids + T, tdm::Q_MAX, r->opt[k], &ansLocal); r->qs[k + 1] = T; r->qe[k + 1] = T + n; for (int j = 0; j < n; j++) { r->pos[T + j] = (uint16_t)(tdm::P_Q + j); r->blk[T + j] = (uint8_t)(k + 1); } for (int i = 0; i < qs[k].n_options; i++) r->opt[k][i] += T; r->ans[k] = T + ansLocal; T += n; } r->T = T; if (allocScratch(mem, T)) break; if (keep <= 1 || keepMax <= 1) { free(seq); free(r); return NO_MEMORY; } } #ifdef TD_TEST g_lastT = r->T < (int)(sizeof(g_lastIds) / sizeof(g_lastIds[0])) ? r->T : (int)(sizeof(g_lastIds) / sizeof(g_lastIds[0])); memcpy(g_lastIds, r->ids, g_lastT * sizeof(uint16_t)); #endif encode(*r, mem); const float* x = (const float*)mem.a; for (int k = 0; k < nq; k++) { Answer& a = out[k]; memset(&a, 0, sizeof(a)); a.type = qs[k].type; if (a.type == CHOICE || a.type == SCORE) headChoice(*r, x, k, qs[k], a); else if (a.type == NOUL) headNoul(*r, x, k, qs[k], a); else headSpan(*r, x, k, state, a); } if (info) { info->tokens = r->T; info->state_tokens = r->S; info->truncated = r->truncated; } freeScratch(mem); free(seq); free(r); if (info) info->ms = now_ms() - t0; #ifdef ESP_PLATFORM ESP_LOGD("tinydecide", "%d questions, %d tokens, %lu ms", nq, info ? info->tokens : 0, (unsigned long)(now_ms() - t0)); #endif return OK; } } // namespace td