Inflate To Deflate
Data is large, and hard drive space is a concern. Real-time effective compression is necessary. Highly evolved compression archivers for text, image, and video trade compression ratio for space and time complexity. Sophisticated context mixing algorithms like PAQ and variants[1] can compress a 1 GB text file into a 116 MB archive. Compression is an AI problem. It is about learning and representing an underlying process that generates the data. The representation is an encoding issue, while the process learning involves techniques ranging from shape coding to pattern segmentation, morphology to neural networks, decision trees to optimization, and matrix factorization to wavelet analysis.
This weekend, I explored a topological approach for data compression. It involves finding an algorithm (decompressor) that yields a tiling topologically equivalent to a boolean matrix. For any file, we get its byte stream. Considering the byte stream as an alias of the original file, we reshape it as a boolean matrix. The dimensions are irrelevant. Depending on the original file size, this matrix will be large (~ billions of rows and columns for a 10 GB file) and dense. To compress such a large binary matrix, we find a [combinatorial] rectangle packing algorithm. Such an algorithm can be found using the chain embedding technique.
The intuition is that the byte stream reshaped as a boolean matrix of dimensions M × N can tell us about its sponginess. An algorithm that generates tiling rules for a boolean matrix satisfying the topological properties is the decompressor, and its embedding into a chain complex becomes the compressed file.
I opted for a new topological approach to compress large boolean matrices because existing methods don't scale well. Transforming it into a packing algorithm discovery problem allows efficient encoding by the decompressor, offering speed without losing information density. Initial experiments show a 25-38 compression ratio on some synthetic images. My focus is on finding the algorithm before fine-tuning for performance, so I've been adjusting inputs for easier discovery.
/*
* A [lossy] non-topological perspective on the same.
*
* A trained network is a lossy encoding of its corpus (can also be boolean sequences): weights and
* activations are finite-precision floats, stored at half precision on disk
* (fp16), and every gradient step discards whatever does not help predict the
* next token. This lossy compressor (?) keeps only the structure that lowers
* cross-entropy and throws the rest away.
*/
#define _GNU_SOURCE
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <math.h>
#include <time.h>
#include <stdint.h>
#include <stdbool.h>
#include <float.h>
#include <getopt.h>
#include <errno.h>
#include <ctype.h>
#include <unistd.h>
#include <pthread.h>
#include <stdatomic.h>
#include <immintrin.h>
#define MAX_VOCAB_SIZE 65536
#define MAX_LAYERS 32
#define EPSILON 1e-6f
#define GRAD_CLIP 50.0f
#define ALIGN64 64
#define MR 16
#define NR 6
#define MAX_MACRO_M 960
#define MAX_MACRO_N 960
#define MAX_MACRO_K 512
#define SMALL_GEMM_THRESHOLD (48 * 48 * 48)
#define PREFETCH(p) _mm_prefetch((const char*)(p), _MM_HINT_T0)
#define BPE_MAX_VOCAB 65536
#define BPE_MAX_PIECE_LEN 128
typedef struct {
int vocab_size;
int n_layer;
int n_embd;
int n_head;
int ctx_len;
int decay_lora_rank;
float ffn_multiplier;
int n_mem_slots;
} lrnnConfig;
static lrnnConfig default_config(void) {
lrnnConfig cfg = {
.vocab_size = 256,
.n_layer = 2,
.n_embd = 64,
.n_head = 1,
.ctx_len = 64,
.decay_lora_rank = 2,
.ffn_multiplier = 2.0f,
.n_mem_slots = 4
};
return cfg;
}
static inline int ffn_hidden(const lrnnConfig *cfg) {
return (int)(cfg->n_embd * cfg->ffn_multiplier);
}
#define MAX_POOL_THREADS 256
typedef void (*task_fn)(void *ctx, long task, int tid);
typedef struct {
task_fn fn;
void *ctx;
long total_tasks;
long job_id;
} pool_job;
static pthread_t g_pool_threads[MAX_POOL_THREADS];
static int g_pool_thread_count = 0;
static bool g_pool_initialized = false;
static bool g_shutdown = false;
static pthread_mutex_t g_mutex = PTHREAD_MUTEX_INITIALIZER;
static pthread_cond_t g_cond_start = PTHREAD_COND_INITIALIZER;
static pthread_cond_t g_cond_done = PTHREAD_COND_INITIALIZER;
static void *xaligned(size_t n);
static pool_job * volatile g_current_job __attribute__((aligned(64))) = NULL;
static atomic_long g_next_task_id __attribute__((aligned(64))) = 0;
static atomic_int g_active_workers __attribute__((aligned(64))) = 0;
static atomic_long g_global_job_id __attribute__((aligned(64))) = 0;
static atomic_int g_job_generation __attribute__((aligned(64))) = 0;
#define SPIN_LIMIT 40000
typedef struct { int tid; } worker_arg;
typedef struct { float *bufA, *bufB; size_t capA, capB; } gemm_scratch;
static gemm_scratch g_scratch[MAX_POOL_THREADS] __attribute__((aligned(64)));
static pthread_mutex_t g_scratch_lock = PTHREAD_MUTEX_INITIALIZER;
static float *scratch_grow(int tid, size_t needA, size_t needB) {
pthread_mutex_lock(&g_scratch_lock);
gemm_scratch *s = &g_scratch[tid];
if (needA > s->capA) { free(s->bufA); s->bufA = (float *)xaligned(needA); s->capA = needA; }
if (needB > s->capB) { free(s->bufB); s->bufB = (float *)xaligned(needB); s->capB = needB; }
pthread_mutex_unlock(&g_scratch_lock);
return s->bufA;
}
static void *pool_worker(void *arg) {
int tid = ((worker_arg *)arg)->tid;
free(arg);
int last_gen = 0;
for (;;) {
long spins = 0;
while (!g_shutdown) {
int gen = atomic_load_explicit(&g_job_generation, memory_order_acquire);
if (gen != last_gen && g_current_job != NULL) break;
if (++spins > SPIN_LIMIT) {
pthread_mutex_lock(&g_mutex);
while (!g_shutdown) {
int gen2 = atomic_load_explicit(&g_job_generation, memory_order_acquire);
if (gen2 != last_gen && g_current_job != NULL) break;
pthread_cond_wait(&g_cond_start, &g_mutex);
}
pthread_mutex_unlock(&g_mutex);
spins = 0;
continue;
}
_mm_pause();
}
if (g_shutdown) break;
last_gen = atomic_load_explicit(&g_job_generation, memory_order_acquire);
pool_job *job = g_current_job;
for (;;) {
long t = atomic_fetch_add(&g_next_task_id, 1);
if (t >= job->total_tasks) break;
job->fn(job->ctx, t, tid);
}
if (atomic_fetch_sub(&g_active_workers, 1) == 1) {
pthread_mutex_lock(&g_mutex);
g_current_job = NULL;
pthread_cond_signal(&g_cond_done);
pthread_mutex_unlock(&g_mutex);
}
}
return NULL;
}
static void pool_init(void) {
pthread_mutex_lock(&g_mutex);
if (g_pool_initialized) { pthread_mutex_unlock(&g_mutex); return; }
int n = (int)sysconf(_SC_NPROCESSORS_ONLN);
if (n < 1) n = 1;
if (n > MAX_POOL_THREADS) n = MAX_POOL_THREADS;
g_pool_thread_count = n;
for (int i = 0; i < n; ++i) {
worker_arg *wa = (worker_arg *)malloc(sizeof(worker_arg));
wa->tid = i;
pthread_create(&g_pool_threads[i], NULL, pool_worker, wa);
}
g_pool_initialized = true;
pthread_mutex_unlock(&g_mutex);
}
static void parallel_for(long total, task_fn fn, void *ctx) {
if (total <= 0) return;
if (!g_pool_initialized) pool_init();
if (total == 1 || g_pool_thread_count == 1) {
for (long t = 0; t < total; ++t) fn(ctx, t, 0);
return;
}
pool_job job;
job.fn = fn; job.ctx = ctx; job.total_tasks = total;
job.job_id = atomic_fetch_add(&g_global_job_id, 1);
atomic_store(&g_next_task_id, 0);
atomic_store(&g_active_workers, g_pool_thread_count);
g_current_job = &job;
atomic_store_explicit(&g_job_generation, job.job_id + 1, memory_order_release);
pthread_cond_broadcast(&g_cond_start);
long spins = 0;
while (g_current_job != NULL) {
if (++spins > SPIN_LIMIT) {
pthread_mutex_lock(&g_mutex);
while (g_current_job != NULL)
pthread_cond_wait(&g_cond_done, &g_mutex);
pthread_mutex_unlock(&g_mutex);
break;
}
_mm_pause();
}
}
static void *xaligned(size_t n) {
void *p = NULL;
if (posix_memalign(&p, ALIGN64, (n + ALIGN64 - 1) & ~(size_t)(ALIGN64 - 1)) != 0 || !p) {
fprintf(stderr, "Error: aligned allocation failed for %zu bytes\n", n);
exit(1);
}
return p;
}
static void transpose_mat(float *dst, const float *src, int rows, int cols);
static void hp_pack_A(long M, long K, const float *A, long lda, float *buf) {
long MP = M / MR, MRr = M % MR, m_mic, k, m;
for (m_mic = 0; m_mic < MP; ++m_mic) {
PREFETCH(A + (m_mic + 1) * MR * lda);
for (m = 0; m < MR; ++m) {
const float *Arow = A + (m_mic * MR + m) * lda;
float *col = buf + m;
for (k = 0; k < K; ++k) { *col = Arow[k]; col += MR; }
}
buf += MR * K;
}
if (MRr > 0) {
for (m = 0; m < MRr; ++m) {
const float *Arow = A + (MP * MR + m) * lda;
float *col = buf + m;
for (k = 0; k < K; ++k) { *col = Arow[k]; col += MR; }
}
for (m = MRr; m < MR; ++m) {
float *col = buf + m;
for (k = 0; k < K; ++k) { *col = 0.0f; col += MR; }
}
}
}
static void hp_pack_B(long K, long N, const float *B, long ldb, float *buf) {
long NP = N / NR, NRr = N % NR, n_mic, k, n;
for (k = 0; k < K; ++k) {
const float *Brow = B + k * ldb;
PREFETCH(B + (k + 1) * ldb);
for (n_mic = 0; n_mic < NP; ++n_mic) {
float *p = buf + n_mic * K * NR + k * NR;
const float *src = Brow + n_mic * NR;
for (n = 0; n < NR; ++n) p[n] = src[n];
}
if (NRr > 0) {
float *p = buf + NP * K * NR + k * NR;
const float *src = Brow + NP * NR;
for (n = 0; n < NRr; ++n) p[n] = src[n];
for (n = NRr; n < NR; ++n) p[n] = 0.0f;
}
}
}
static inline __attribute__((always_inline)) void
micro_16x6(long kc, const float *A, const float *B,
float *C, long ldc, int mc, int nc, int beta_is_zero)
{
float AB[MR * NR] __attribute__((aligned(ALIGN64)));
long k_main = kc / 4, k_rem = kc % 4;
const float *a = A, *b = B; float *ab = AB;
__asm__ __volatile__(
"vxorps %%ymm0,%%ymm0,%%ymm0 \n\t vxorps %%ymm1,%%ymm1,%%ymm1 \n\t"
"vxorps %%ymm2,%%ymm2,%%ymm2 \n\t vxorps %%ymm3,%%ymm3,%%ymm3 \n\t"
"vxorps %%ymm4,%%ymm4,%%ymm4 \n\t vxorps %%ymm5,%%ymm5,%%ymm5 \n\t"
"vxorps %%ymm6,%%ymm6,%%ymm6 \n\t vxorps %%ymm7,%%ymm7,%%ymm7 \n\t"
"vxorps %%ymm8,%%ymm8,%%ymm8 \n\t vxorps %%ymm9,%%ymm9,%%ymm9 \n\t"
"vxorps %%ymm10,%%ymm10,%%ymm10 \n\t vxorps %%ymm11,%%ymm11,%%ymm11 \n\t"
"test %[km],%[km] \n\t je 2f \n\t"
"1: \n\t"
#define KSTEP(AO, BO) \
"vmovaps " #AO "(%[a]),%%ymm12 \n\t" \
"vmovaps " #AO "+32(%[a]),%%ymm13 \n\t" \
"vbroadcastss " #BO "(%[b]),%%ymm14 \n\t" \
"vfmadd231ps %%ymm12,%%ymm14,%%ymm0 \n\t" \
"vfmadd231ps %%ymm13,%%ymm14,%%ymm1 \n\t" \
"vbroadcastss " #BO "+4(%[b]),%%ymm15 \n\t" \
"vfmadd231ps %%ymm12,%%ymm15,%%ymm2 \n\t" \
"vfmadd231ps %%ymm13,%%ymm15,%%ymm3 \n\t" \
"vbroadcastss " #BO "+8(%[b]),%%ymm14 \n\t" \
"vfmadd231ps %%ymm12,%%ymm14,%%ymm4 \n\t" \
"vfmadd231ps %%ymm13,%%ymm14,%%ymm5 \n\t" \
"vbroadcastss " #BO "+12(%[b]),%%ymm15 \n\t" \
"vfmadd231ps %%ymm12,%%ymm15,%%ymm6 \n\t" \
"vfmadd231ps %%ymm13,%%ymm15,%%ymm7 \n\t" \
"vbroadcastss " #BO "+16(%[b]),%%ymm14 \n\t" \
"vfmadd231ps %%ymm12,%%ymm14,%%ymm8 \n\t" \
"vfmadd231ps %%ymm13,%%ymm14,%%ymm9 \n\t" \
"vbroadcastss " #BO "+20(%[b]),%%ymm15 \n\t" \
"vfmadd231ps %%ymm12,%%ymm15,%%ymm10 \n\t" \
"vfmadd231ps %%ymm13,%%ymm15,%%ymm11 \n\t"
KSTEP(0, 0)
KSTEP(64, 24)
KSTEP(128, 48)
KSTEP(192, 72)
#undef KSTEP
"add $256,%[a] \n\t add $96,%[b] \n\t"
"dec %[km] \n\t jnz 1b \n\t"
"2: \n\t"
"test %[kr],%[kr] \n\t je 4f \n\t"
"3: \n\t"
"vmovaps 0(%[a]),%%ymm12 \n\t vmovaps 32(%[a]),%%ymm13 \n\t"
"add $64,%[a] \n\t"
"vbroadcastss 0(%[b]),%%ymm14 \n\t"
"vfmadd231ps %%ymm12,%%ymm14,%%ymm0 \n\t vfmadd231ps %%ymm13,%%ymm14,%%ymm1 \n\t"
"vbroadcastss 4(%[b]),%%ymm15 \n\t"
"vfmadd231ps %%ymm12,%%ymm15,%%ymm2 \n\t vfmadd231ps %%ymm13,%%ymm15,%%ymm3 \n\t"
"vbroadcastss 8(%[b]),%%ymm14 \n\t"
"vfmadd231ps %%ymm12,%%ymm14,%%ymm4 \n\t vfmadd231ps %%ymm13,%%ymm14,%%ymm5 \n\t"
"vbroadcastss 12(%[b]),%%ymm15 \n\t"
"vfmadd231ps %%ymm12,%%ymm15,%%ymm6 \n\t vfmadd231ps %%ymm13,%%ymm15,%%ymm7 \n\t"
"vbroadcastss 16(%[b]),%%ymm14 \n\t"
"vfmadd231ps %%ymm12,%%ymm14,%%ymm8 \n\t vfmadd231ps %%ymm13,%%ymm14,%%ymm9 \n\t"
"vbroadcastss 20(%[b]),%%ymm15 \n\t"
"vfmadd231ps %%ymm12,%%ymm15,%%ymm10 \n\t vfmadd231ps %%ymm13,%%ymm15,%%ymm11 \n\t"
"add $24,%[b] \n\t dec %[kr] \n\t jnz 3b \n\t"
"4: \n\t"
"vmovups %%ymm0,0(%[ab]) \n\t vmovups %%ymm1,32(%[ab]) \n\t"
"vmovups %%ymm2,64(%[ab]) \n\t vmovups %%ymm3,96(%[ab]) \n\t"
"vmovups %%ymm4,128(%[ab]) \n\t vmovups %%ymm5,160(%[ab]) \n\t"
"vmovups %%ymm6,192(%[ab]) \n\t vmovups %%ymm7,224(%[ab]) \n\t"
"vmovups %%ymm8,256(%[ab]) \n\t vmovups %%ymm9,288(%[ab]) \n\t"
"vmovups %%ymm10,320(%[ab]) \n\t vmovups %%ymm11,352(%[ab]) \n\t"
"vzeroupper \n\t"
: [km] "+r"(k_main), [kr] "+r"(k_rem), [a] "+r"(a), [b] "+r"(b)
: [ab] "r"(ab)
: "memory", "ymm0","ymm1","ymm2","ymm3","ymm4","ymm5","ymm6","ymm7",
"ymm8","ymm9","ymm10","ymm11","ymm12","ymm13","ymm14","ymm15");
if (beta_is_zero)
for (int i = 0; i < mc; i++)
for (int c = 0; c < nc; c++)
C[(long)i * ldc + c] = AB[i + c * MR];
else
for (int i = 0; i < mc; i++)
for (int c = 0; c < nc; c++)
C[(long)i * ldc + c] += AB[i + c * MR];
}
static void hp_macro_kernel(long M, long N, long K, const float *pA, const float *pB,
float *C, long ldc, int beta_is_zero) {
long MP = (M + MR - 1) / MR, NP = (N + NR - 1) / NR;
int Mr = (int)(M % MR), Nr = (int)(N % NR);
for (long im = 0; im < MP; ++im) {
int mc = (im == MP - 1 && Mr) ? Mr : MR;
const float *A0 = pA + im * MR * K;
for (long jn = 0; jn < NP; ++jn) {
int nc = (jn == NP - 1 && Nr) ? Nr : NR;
const float *B0 = pB + jn * K * NR;
float *C0 = C + im * MR * ldc + jn * NR;
micro_16x6(K, A0, B0, C0, ldc, mc, nc, beta_is_zero);
}
}
}
typedef struct {
const float *A;
const float *B;
float *C;
long M, N, K;
long lda, ldb, ldc;
long blk_m, blk_n, blk_k;
long num_m_tiles, num_n_tiles;
float **bufA;
float **bufB;
} hp_gemm_ctx;
static void hp_gemm_tile(void *vctx, long task, int tid) {
hp_gemm_ctx *g = (hp_gemm_ctx *)vctx;
long it = task / g->num_n_tiles, jt = task % g->num_n_tiles;
long ic = it * g->blk_m, jc = jt * g->blk_n;
long mc = g->M - ic; if (mc > g->blk_m) mc = g->blk_m;
long nc = g->N - jc; if (nc > g->blk_n) nc = g->blk_n;
float *bufA = g->bufA[tid], *bufB = g->bufB[tid];
float *Ctile = g->C + ic * g->ldc + jc;
for (long pc = 0; pc < g->K; pc += g->blk_k) {
long kc = g->K - pc; if (kc > g->blk_k) kc = g->blk_k;
hp_pack_A(mc, kc, g->A + ic * g->lda + pc, g->lda, bufA);
hp_pack_B(kc, nc, g->B + pc * g->ldb + jc, g->ldb, bufB);
hp_macro_kernel(mc, nc, kc, bufA, bufB, Ctile, g->ldc, pc == 0);
}
}
typedef struct { const float *x; const float *B; float *y; long N, K, nblk; } hp_gemv_ctx;
static void hp_gemv_block(const float *x, const float *B, float *y, long N, long K, long n0, long n1) {
for (long n = n0; n < n1; ++n) y[n] = 0.0f;
for (long k = 0; k < K; ++k) {
const float *Brow = B + k * N;
__m256 xk = _mm256_set1_ps(x[k]);
long n = n0;
for (; n + 8 <= n1; n += 8) {
__m256 w = _mm256_loadu_ps(Brow + n);
_mm256_storeu_ps(y + n, _mm256_fmadd_ps(xk, w, _mm256_loadu_ps(y + n)));
}
for (; n < n1; ++n) y[n] += x[k] * Brow[n];
}
}
static void hp_gemv_task(void *vc, long t, int tid) {
(void)tid;
hp_gemv_ctx *c = (hp_gemv_ctx *)vc;
long per = (c->N + c->nblk - 1) / c->nblk;
long n0 = t * per, n1 = n0 + per; if (n1 > c->N) n1 = c->N;
if (n0 < n1) hp_gemv_block(c->x, c->B, c->y, c->N, c->K, n0, n1);
}
static void hp_sgemm(float *C, const float *A, const float *B,
long M, long N, long K) {
if (M == 1) {
if (N * K < 600000) {
hp_gemv_block(A, B, C, N, K, 0, N);
} else {
if (!g_pool_initialized) pool_init();
long nblk = g_pool_thread_count;
hp_gemv_ctx c = { A, B, C, N, K, nblk };
parallel_for(nblk, hp_gemv_task, &c);
}
return;
}
hp_gemm_ctx g;
g.A = A; g.B = B; g.C = C; g.M = M; g.N = N; g.K = K;
g.lda = K; g.ldb = N; g.ldc = N;
unsigned long long sz = (unsigned long long)M * N * K;
long bm = (sz < (unsigned long long)2000 * 2000) ? MAX_MACRO_M : 240;
long bn = (sz < (unsigned long long)2000 * 2000) ? MAX_MACRO_N : 240;
if (bm > M) bm = M;
if (bn > N) bn = N;
long bk = K; if (bk > MAX_MACRO_K) bk = MAX_MACRO_K;
g.blk_m = bm; g.blk_n = bn; g.blk_k = bk;
g.num_m_tiles = (M + bm - 1) / bm;
g.num_n_tiles = (N + bn - 1) / bn;
long total = g.num_m_tiles * g.num_n_tiles;
int nthr = g_pool_initialized ? g_pool_thread_count : 1;
if (!g_pool_initialized && sz >= SMALL_GEMM_THRESHOLD) { pool_init(); nthr = g_pool_thread_count; }
long padA = ((bm + MR - 1) / MR) * MR;
long padB = ((bn + NR - 1) / NR) * NR;
size_t aBytes = (size_t)padA * bk * sizeof(float);
size_t bBytes = (size_t)bk * padB * sizeof(float);
int use = (total == 1 || sz < SMALL_GEMM_THRESHOLD) ? 1 : nthr;
if (use > MAX_POOL_THREADS) use = MAX_POOL_THREADS;
float **bufA = (float **)alloca((size_t)use * sizeof(float *));
float **bufB = (float **)alloca((size_t)use * sizeof(float *));
for (int i = 0; i < use; ++i) {
bufA[i] = scratch_grow(i, aBytes, bBytes);
bufB[i] = g_scratch[i].bufB;
}
g.bufA = bufA; g.bufB = bufB;
if (use == 1) { for (long t = 0; t < total; ++t) hp_gemm_tile(&g, t, 0); }
else parallel_for(total, hp_gemm_tile, &g);
}
typedef struct {
float *data;
int rows;
int cols;
int size;
} Tensor;
static Tensor tensor_alloc(int rows, int cols) {
Tensor t;
t.rows = rows;
t.cols = cols;
t.size = rows * cols;
t.data = NULL;
if (t.size > 0) {
t.data = (float *)calloc((size_t)t.size, sizeof(float));
if (!t.data) {
fprintf(stderr, "Error: tensor allocation failed (%d x %d)\n", rows, cols);
exit(1);
}
}
return t;
}
static Tensor tensor_alloc_1d(int size) {
return tensor_alloc(size, 1);
}
static void tensor_free(Tensor *t) {
if (t && t->data) {
free(t->data);
t->data = NULL;
}
if (t) {
t->rows = t->cols = t->size = 0;
}
}
static void tensor_copy(Tensor *dst, const Tensor *src) {
if (dst->size != src->size) {
fprintf(stderr, "Error: tensor_copy size mismatch (%d vs %d)\n", dst->size, src->size);
exit(1);
}
if (src->size > 0) {
memcpy(dst->data, src->data, (size_t)src->size * sizeof(float));
}
}
static void tensor_fill(Tensor *t, float val) {
for (int i = 0; i < t->size; i++) {
t->data[i] = val;
}
}
static void tensor_zero(Tensor *t) {
if (t->data && t->size > 0) {
memset(t->data, 0, (size_t)t->size * sizeof(float));
}
}
static float randn(void) {
float u1 = ((float)rand() + 1.0f) / ((float)RAND_MAX + 2.0f);
float u2 = ((float)rand() + 1.0f) / ((float)RAND_MAX + 2.0f);
return sqrtf(-2.0f * logf(u1)) * cosf(2.0f * 3.14159265f * u2);
}
static float rand_uniform(float lo, float hi) {
return lo + ((float)rand() / (float)RAND_MAX) * (hi - lo);
}
static void tensor_randn(Tensor *t, float scale) {
for (int i = 0; i < t->size; i++) {
t->data[i] = randn() * scale;
}
}
static void tensor_rand_uniform(Tensor *t, float lo, float hi) {
for (int i = 0; i < t->size; i++) {
t->data[i] = rand_uniform(lo, hi);
}
}
static inline float sigmoid_f(float x) {
return 1.0f / (1.0f + expf(-x));
}
static inline float silu_f(float x) {
return x * sigmoid_f(x);
}
static inline float clamp_f(float x, float lo, float hi) {
if (x < lo) return lo;
if (x > hi) return hi;
return x;
}
static inline __m256 exp256_ps(__m256 x) {
const __m256 p0 = _mm256_set1_ps(1.5400435E-4f);
const __m256 p1 = _mm256_set1_ps(1.3329787E-3f);
const __m256 p2 = _mm256_set1_ps(9.6151903E-3f);
const __m256 p3 = _mm256_set1_ps(5.5548714E-2f);
const __m256 p4 = _mm256_set1_ps(2.40226507E-1f);
const __m256 p5 = _mm256_set1_ps(6.93147181E-1f);
const __m256 l2e = _mm256_set1_ps(1.442695040f);
const __m256 magic = _mm256_set1_ps(12582912.0f);
x = _mm256_mul_ps(x, l2e);
__m256 t = _mm256_add_ps(x, magic);
__m256 n = _mm256_sub_ps(t, magic);
__m256 f = _mm256_sub_ps(x, n);
__m256 p = p0;
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, p1);
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, p2);
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, p3);
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, p4);
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, p5);
p = _mm256_mul_ps(p, f);
p = _mm256_add_ps(p, _mm256_set1_ps(1.0f));
__m256i ni = _mm256_cvtps_epi32(n);
__m256 scale = _mm256_castsi256_ps(_mm256_slli_epi32(_mm256_add_epi32(ni, _mm256_set1_epi32(127)), 23));
return _mm256_mul_ps(p, scale);
}
static inline float exp_fast(float x) {
float r __attribute__((aligned(32)));
_mm256_store_ps(&r, exp256_ps(_mm256_set1_ps(x)));
return r;
}
static void sigmoid_vec(float *out, const float *in, int n) {
int i = 0;
for (; i + 8 <= n; i += 8) {
__m256 x = _mm256_loadu_ps(in + i);
__m256 s = exp256_ps(_mm256_sub_ps(_mm256_setzero_ps(), x));
_mm256_storeu_ps(out + i, _mm256_div_ps(_mm256_set1_ps(1.0f), _mm256_add_ps(_mm256_set1_ps(1.0f), s)));
}
for (; i < n; i++) out[i] = sigmoid_f(in[i]);
}
static void exp_vec(float *out, const float *in, int n) {
__m256 lo = _mm256_set1_ps(-10.0f), hi = _mm256_set1_ps(10.0f);
int i = 0;
for (; i + 8 <= n; i += 8) {
__m256 x = _mm256_loadu_ps(in + i);
x = _mm256_max_ps(lo, _mm256_min_ps(x, hi));
_mm256_storeu_ps(out + i, exp256_ps(x));
}
for (; i < n; i++) out[i] = expf(clamp_f(in[i], -10.0f, 10.0f));
}
static void softmax_vec(float *out, const float *in, int n) {
float max_val = in[0];
for (int i = 1; i < n; i++) {
if (in[i] > max_val) max_val = in[i];
}
int i = 0;
__m256 vmax = _mm256_set1_ps(max_val), vs = _mm256_setzero_ps();
float sum = 0.0f;
if (n >= 16) {
for (; i + 8 <= n; i += 8) {
__m256 e = exp256_ps(_mm256_sub_ps(_mm256_loadu_ps(in + i), vmax));
_mm256_storeu_ps(out + i, e);
vs = _mm256_add_ps(vs, e);
}
__m128 lo = _mm_add_ps(_mm256_castps256_ps128(vs), _mm256_extractf128_ps(vs, 1));
lo = _mm_hadd_ps(lo, lo); lo = _mm_hadd_ps(lo, lo);
sum = _mm_cvtss_f32(lo);
}
for (; i < n; i++) {
out[i] = exp_fast(in[i] - max_val);
sum += out[i];
}
float inv_sum = 1.0f / (sum + EPSILON);
i = 0;
if (n >= 16) {
__m256 vis = _mm256_set1_ps(inv_sum);
for (; i + 8 <= n; i += 8)
_mm256_storeu_ps(out + i, _mm256_mul_ps(_mm256_loadu_ps(out + i), vis));
}
for (; i < n; i++) out[i] *= inv_sum;
}
static inline __m256 clamp256(__m256 x) {
return _mm256_max_ps(_mm256_set1_ps(-10.0f), _mm256_min_ps(x, _mm256_set1_ps(10.0f)));
}
static inline __m256 sigmoid256(__m256 x) {
__m256 s = exp256_ps(_mm256_sub_ps(_mm256_setzero_ps(), x));
return _mm256_div_ps(_mm256_set1_ps(1.0f), _mm256_add_ps(_mm256_set1_ps(1.0f), s));
}
static void vec_add(float *out, const float *a, const float *b, int n) {
int i = 0;
for (; i + 8 <= n; i += 8) {
__m256 va = _mm256_loadu_ps(a + i);
__m256 vb = _mm256_loadu_ps(b + i);
_mm256_storeu_ps(out + i, _mm256_add_ps(va, vb));
}
for (; i < n; i++) {
out[i] = a[i] + b[i];
}
}
static void vec_mul(float *out, const float *a, const float *b, int n) {
int i = 0;
for (; i + 8 <= n; i += 8) {
__m256 va = _mm256_loadu_ps(a + i);
__m256 vb = _mm256_loadu_ps(b + i);
_mm256_storeu_ps(out + i, _mm256_mul_ps(va, vb));
}
for (; i < n; i++) {
out[i] = a[i] * b[i];
}
}
static void matvec(float *out, const float *x, const Tensor *W) {
hp_sgemm(out, x, W->data, 1, W->cols, W->rows);
}
static void matmul(Tensor *out, const Tensor *X, const Tensor *W) {
if (W->rows != X->cols) {
fprintf(stderr, "matmul dimension mismatch: X(%d,%d) @ W(%d,%d)\n",
X->rows, X->cols, W->rows, W->cols);
exit(1);
}
hp_sgemm(out->data, X->data, W->data, X->rows, W->cols, X->cols);
}
static void layer_norm(float *out, const float *x, const float *weight,
const float *bias, int n) {
__m256 acc_mean = _mm256_setzero_ps();
int i = 0;
for (; i + 8 <= n; i += 8) {
acc_mean = _mm256_add_ps(acc_mean, _mm256_loadu_ps(x + i));
}
__m128 lo = _mm_add_ps(_mm256_castps256_ps128(acc_mean), _mm256_extractf128_ps(acc_mean, 1));
lo = _mm_hadd_ps(lo, lo); lo = _mm_hadd_ps(lo, lo);
float mean = _mm_cvtss_f32(lo);
for (; i < n; i++) mean += x[i];
mean /= (float)n;
__m256 vmean = _mm256_set1_ps(mean);
__m256 acc_var = _mm256_setzero_ps();
i = 0;
for (; i + 8 <= n; i += 8) {
__m256 d = _mm256_sub_ps(_mm256_loadu_ps(x + i), vmean);
acc_var = _mm256_fmadd_ps(d, d, acc_var);
}
lo = _mm_add_ps(_mm256_castps256_ps128(acc_var), _mm256_extractf128_ps(acc_var, 1));
lo = _mm_hadd_ps(lo, lo); lo = _mm_hadd_ps(lo, lo);
float var = _mm_cvtss_f32(lo);
for (; i < n; i++) {
float d = x[i] - mean;
var += d * d;
}
var /= (float)n;
float inv_std = 1.0f / sqrtf(var + EPSILON);
__m256 vinv_std = _mm256_set1_ps(inv_std);
i = 0;
for (; i + 8 <= n; i += 8) {
__m256 vx = _mm256_loadu_ps(x + i);
__m256 vw = _mm256_loadu_ps(weight + i);
__m256 vb = _mm256_loadu_ps(bias + i);
__m256 normalized = _mm256_mul_ps(_mm256_sub_ps(vx, vmean), vinv_std);
_mm256_storeu_ps(out + i, _mm256_fmadd_ps(vw, normalized, vb));
}
for (; i < n; i++) {
out[i] = weight[i] * (x[i] - mean) * inv_std + bias[i];
}
}
static void layer_norm_seq(Tensor *out, const Tensor *x, const Tensor *weight,
const Tensor *bias) {
int seq_len = x->rows;
int n_embd = x->cols;
for (int s = 0; s < seq_len; s++) {
layer_norm(out->data + s * n_embd,
x->data + s * n_embd,
weight->data, bias->data, n_embd);
}
}
static void compute_alibi_slopes(float *slopes, int n_head) {
for (int h = 0; h < n_head; h++) {
float exponent = -8.0f * (float)(h + 1) / (float)n_head;
slopes[h] = powf(2.0f, exponent);
}
}
typedef struct {
Tensor ln1_weight, ln1_bias;
Tensor ln2_weight, ln2_bias;
Tensor time_shift_w1, time_shift_w2, time_shift_w4;
Tensor time_mix_r, time_mix_k, time_mix_v;
Tensor decay_lora_a, decay_lora_b;
Tensor decay_base;
Tensor time_first;
Tensor Wr, Wk, Wv, Wo;
Tensor channel_mix;
Tensor ffn_gate_up;
Tensor ffn_down;
Tensor alibi_slopes;
Tensor mem_gate_write;
Tensor mem_gate_read;
} LayerParams;
typedef struct {
Tensor emb;
Tensor ln0_weight, ln0_bias;
LayerParams *layers;
Tensor ln_out_weight, ln_out_bias;
Tensor head;
int n_layers;
} ModelParams;
typedef struct {
Tensor x_prev_1, x_prev_2, x_prev_3, x_prev_4;
Tensor wkv_num;
Tensor wkv_den;
Tensor ffn_prev;
} LayerState;
typedef struct {
LayerState *layers;
int n_layers;
float *scr;
int scr_n_embd, scr_ffn_h, scr_n_slots;
} ModelState;
typedef struct {
char **piece_str;
int *piece_len;
float *piece_score;
int8_t *piece_type;
int byte_tok[256];
int *hash_id;
int hash_mask;
int vocab_size;
int unk_id;
int pad_id;
} BPEVocab;
static uint64_t fnv1a(const char *s, int n) {
uint64_t h = 1469598103934665603ULL;
for (int i = 0; i < n; ++i) { h ^= (unsigned char)s[i]; h *= 1099511628211ULL; }
return h;
}
static int bpe_lookup(const BPEVocab *bpe, const char *s, int n) {
uint64_t h = fnv1a(s, n) & (uint64_t)bpe->hash_mask;
for (;;) {
int id = bpe->hash_id[h];
if (id < 0) return -1;
if (bpe->piece_len[id] == n && memcmp(bpe->piece_str[id], s, (size_t)n) == 0) return id;
h = (h + 1) & (uint64_t)bpe->hash_mask;
}
}
typedef struct { int start, end, prev, next, active, id; } SymTok;
typedef struct { float score; int left, right; } BPECand;
typedef struct { BPECand *a; int n, cap; } BPEHeap;
static void bpe_heap_push(BPEHeap *H, BPECand c) {
if (H->n == H->cap) { H->cap = H->cap ? H->cap * 2 : 64; H->a = (BPECand *)realloc(H->a, (size_t)H->cap * sizeof(BPECand)); }
int i = H->n++; H->a[i] = c;
while (i > 0) {
int p = (i - 1) / 2;
if (H->a[p].score > H->a[i].score ||
(H->a[p].score == H->a[i].score && H->a[p].left <= H->a[i].left)) break;
BPECand t = H->a[p]; H->a[p] = H->a[i]; H->a[i] = t; i = p;
}
}
static int bpe_heap_pop(BPEHeap *H, BPECand *out) {
if (H->n == 0) return 0;
*out = H->a[0]; H->a[0] = H->a[--H->n];
int i = 0;
for (;;) {
int l = 2 * i + 1, r = 2 * i + 2, b = i;
if (l < H->n && (H->a[l].score > H->a[b].score ||
(H->a[l].score == H->a[b].score && H->a[l].left < H->a[b].left))) b = l;
if (r < H->n && (H->a[r].score > H->a[b].score ||
(H->a[r].score == H->a[b].score && H->a[r].left < H->a[b].left))) b = r;
if (b == i) break;
BPECand t = H->a[b]; H->a[b] = H->a[i]; H->a[i] = t; i = b;
}
return 1;
}
static int utf8_len(unsigned char c) {
if (c < 0x80) return 1;
if ((c >> 5) == 0x6) return 2;
if ((c >> 4) == 0xE) return 3;
if ((c >> 3) == 0x1E) return 4;
return 1;
}
typedef struct { uint64_t key; int count; int used; } BpeBigramEntry;
typedef struct { int64_t key; int count; } BpeHeapEnt;
static uint64_t bpe_pair_key(int a, int b) { return ((uint64_t)(uint32_t)a << 32) | (uint32_t)b; }
static void bpe_merge_heap_push(BpeHeapEnt **heap, int *heap_n, int *heap_cap, int count, int a, int b) {
if (*heap_n >= *heap_cap) {
*heap_cap *= 2;
*heap = (BpeHeapEnt *)realloc(*heap, (size_t)*heap_cap * sizeof(BpeHeapEnt));
}
int i = (*heap_n)++;
(*heap)[i].count = count; (*heap)[i].key = (int64_t)bpe_pair_key(a, b);
while (i > 0) {
int p = (i - 1) / 2;
if ((*heap)[p].count > (*heap)[i].count ||
((*heap)[p].count == (*heap)[i].count && (*heap)[p].key <= (*heap)[i].key)) break;
BpeHeapEnt t = (*heap)[p]; (*heap)[p] = (*heap)[i]; (*heap)[i] = t; i = p;
}
}
static int bpe_merge_heap_pop(BpeHeapEnt *heap, int *heap_n, int *out_a, int *out_b) {
if (*heap_n == 0) return 0;
int64_t key = heap[0].key; int cnt = heap[0].count;
heap[0] = heap[--(*heap_n)];
int i = 0;
for (;;) {
int l = 2 * i + 1, r = 2 * i + 2, b = i;
if (l < *heap_n && (heap[l].count > heap[b].count ||
(heap[l].count == heap[b].count && heap[l].key < heap[b].key))) b = l;
if (r < *heap_n && (heap[r].count > heap[b].count ||
(heap[r].count == heap[b].count && heap[r].key < heap[b].key))) b = r;
if (b == i) break;
BpeHeapEnt t = heap[b]; heap[b] = heap[i]; heap[i] = t; i = b;
}
*out_a = (int)(key >> 32); *out_b = (int)(key & 0xFFFFFFFFu);
return cnt;
}
static void bpe_map_set(BpeHeapEnt **heap, int *heap_n, int *heap_cap,
BpeBigramEntry *bigrams, int mask, int a, int b, int delta) {
uint64_t key = bpe_pair_key(a, b);
uint64_t slot = key & (uint64_t)mask;
while (bigrams[slot].used && bigrams[slot].key != key) slot = (slot + 1) & (uint64_t)mask;
if (!bigrams[slot].used) { bigrams[slot].used = 1; bigrams[slot].key = key; bigrams[slot].count = 0; }
bigrams[slot].count += delta;
if (bigrams[slot].count < 0) bigrams[slot].count = 0;
if (bigrams[slot].count > 0) bpe_merge_heap_push(heap, heap_n, heap_cap, bigrams[slot].count, a, b);
}
static int bpe_map_get(const BpeBigramEntry *bigrams, int mask, int a, int b) {
uint64_t key = bpe_pair_key(a, b);
uint64_t slot = key & (uint64_t)mask;
while (bigrams[slot].used && bigrams[slot].key != key) slot = (slot + 1) & (uint64_t)mask;
return bigrams[slot].used ? bigrams[slot].count : 0;
}
static int bpe_pair_in_vocab(const BPEVocab *b, int a, int c) {
if (a < 0 || c < 0) return 1;
const char *sa = b->piece_str[a], *sb = b->piece_str[c];
int la = b->piece_len[a], lb = b->piece_len[c];
if (la + lb > BPE_MAX_PIECE_LEN * 4) return 1;
char tmp[BPE_MAX_PIECE_LEN * 4 + 1];
memcpy(tmp, sa, (size_t)la); memcpy(tmp + la, sb, (size_t)lb);
return bpe_lookup(b, tmp, la + lb) >= 0;
}
static void bpe_build_vocab(BPEVocab *bpe, const char *text, size_t text_len, int target_vocab_size) {
if (target_vocab_size > BPE_MAX_VOCAB) target_vocab_size = BPE_MAX_VOCAB;
if (target_vocab_size < 256 + 16) target_vocab_size = 256 + 16;
int capacity = target_vocab_size + 16;
bpe->piece_str = (char **)calloc((size_t)capacity, sizeof(char *));
bpe->piece_len = (int *)calloc((size_t)capacity, sizeof(int));
bpe->piece_score = (float *)calloc((size_t)capacity, sizeof(float));
bpe->piece_type = (int8_t *)calloc((size_t)capacity, sizeof(int8_t));
for (int i = 0; i < 256; i++) bpe->byte_tok[i] = -1;
int vocab_count = 0;
void add_piece(const char *s, int type, float score) {
int len = (int)strlen(s);
bpe->piece_str[vocab_count] = (char *)malloc((size_t)(len + 1));
memcpy(bpe->piece_str[vocab_count], s, (size_t)(len + 1));
bpe->piece_len[vocab_count] = len;
bpe->piece_score[vocab_count] = score;
bpe->piece_type[vocab_count] = (int8_t)type;
vocab_count++;
}
add_piece("", 1, 0.0f); bpe->unk_id = 0;
add_piece("", 2, 0.0f);
add_piece("", 2, 0.0f);
bpe->pad_id = -1;
for (int b = 0; b < 256; b++) {
char buf[8];
snprintf(buf, sizeof(buf), "<0x%02X>", b);
add_piece(buf, 3, 0.0f);
bpe->byte_tok[b] = vocab_count - 1;
}
bool seen_chars[256] = {false};
for (size_t i = 0; i < text_len; ) {
int clen = utf8_len((unsigned char)text[i]);
if (i + clen > text_len) break;
if (clen == 1 && !seen_chars[(unsigned char)text[i]]) {
seen_chars[(unsigned char)text[i]] = true;
char buf[2] = { text[i], '\0' };
if (vocab_count < capacity) {
add_piece(buf, 0, -1e6f);
}
}
if (i == 0 || text[i-1] == ' ' || text[i-1] == '\n') {
char buf[8] = {(char)0xE2, (char)0x96, (char)0x81, 0};
int blen = 3;
for (int j = 0; j < clen && blen < 7; j++) {
buf[blen++] = text[i + j];
}
buf[blen] = '\0';
bool found = false;
for (int k = 0; k < vocab_count; k++) {
if (bpe->piece_len[k] == blen && memcmp(bpe->piece_str[k], buf, (size_t)blen) == 0) {
found = true;
break;
}
}
if (!found && vocab_count < capacity) {
add_piece(buf, 0, -1e6f);
}
}
i += (size_t)clen;
}
{
const char *uw = "\xE2\x96\x81";
bool found = false;
for (int k = 0; k < vocab_count; k++) {
if (bpe->piece_len[k] == 3 && memcmp(bpe->piece_str[k], uw, 3) == 0) { found = true; break; }
}
if (!found && vocab_count < capacity) add_piece(uw, 0, -1e6f);
}
size_t max_syms = text_len * 2 + 16;
SymTok *syms = (SymTok *)malloc(sizeof(SymTok) * max_syms);
char *buf = (char *)malloc(text_len * 3 + 4);
int bn = 0;
buf[bn++] = (char)0xE2; buf[bn++] = (char)0x96; buf[bn++] = (char)0x81;
for (size_t i = 0; i < text_len; i++) {
if (text[i] == ' ') {
buf[bn++] = (char)0xE2; buf[bn++] = (char)0x96; buf[bn++] = (char)0x81;
} else {
buf[bn++] = text[i];
}
}
int ns = 0;
for (int i = 0; i < bn; ) {
int l = utf8_len((unsigned char)buf[i]);
if (i + l > bn) l = bn - i;
if ((size_t)ns >= max_syms - 1) break;
syms[ns].start = i; syms[ns].end = i + l;
syms[ns].prev = ns - 1; syms[ns].next = ns + 1; syms[ns].active = 1;
syms[ns].id = -1;
if ((unsigned char)buf[i] == 0xE2 && i + 3 <= bn && (unsigned char)buf[i+1] == 0x96 && (unsigned char)buf[i+2] == 0x81) {
if (i + l < bn) {
int cl2 = utf8_len((unsigned char)buf[i + l]);
if (i + l + cl2 <= bn) {
syms[ns].end = i + l + cl2;
l = syms[ns].end - i;
}
}
}
ns++; i += l;
}
if (ns > 0) syms[ns - 1].next = -1;
int hs = 1;
while (hs < target_vocab_size * 2 + 32) hs <<= 1;
bpe->hash_mask = hs - 1;
bpe->hash_id = (int *)malloc((size_t)hs * sizeof(int));
for (int i = 0; i < hs; i++) bpe->hash_id[i] = -1;
for (int k = 1; k < vocab_count; k++) {
if (bpe->piece_type[k] != 1 && bpe->piece_type[k] != 2) {
uint64_t h = fnv1a(bpe->piece_str[k], bpe->piece_len[k]) & (uint64_t)bpe->hash_mask;
uint64_t slot = h & (uint64_t)bpe->hash_mask;
while (bpe->hash_id[slot] >= 0) slot = (slot + 1) & (uint64_t)bpe->hash_mask;
bpe->hash_id[slot] = k;
}
}
int map_cap = 16;
while (map_cap < ns * 8) map_cap <<= 1;
BpeBigramEntry *bigrams = (BpeBigramEntry *)calloc((size_t)map_cap, sizeof(BpeBigramEntry));
int bigram_mask = map_cap - 1;
int heap_cap = (ns * 4) + 16;
BpeHeapEnt *heap = (BpeHeapEnt *)malloc(sizeof(BpeHeapEnt) * (size_t)heap_cap);
int heap_n = 0;
for (int i = 0; i < ns - 1; i++) {
int a = syms[i].id = bpe_lookup(bpe, buf + syms[i].start, syms[i].end - syms[i].start);
int b = syms[i + 1].id = bpe_lookup(bpe, buf + syms[i + 1].start, syms[i + 1].end - syms[i + 1].start);
if (a >= 0 && b >= 0 && !bpe_pair_in_vocab(bpe, a, b)) bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, a, b, 1);
}
if (ns > 0) syms[ns - 1].id = bpe_lookup(bpe, buf + syms[ns - 1].start, syms[ns - 1].end - syms[ns - 1].start);
int merges_done = 0;
int max_merges = target_vocab_size - vocab_count;
if (max_merges < 0) max_merges = 0;
while (merges_done < max_merges) {
int la, ra, cnt;
do {
cnt = bpe_merge_heap_pop(heap, &heap_n, &la, &ra);
if (cnt == 0) break;
} while (bpe_map_get(bigrams, bigram_mask, la, ra) != cnt);
if (cnt == 0) break;
int la_len = bpe->piece_len[la], ra_len = bpe->piece_len[ra];
if (vocab_count >= capacity) break;
char *np = (char *)malloc((size_t)(la_len + ra_len + 1));
memcpy(np, bpe->piece_str[la], (size_t)la_len);
memcpy(np + la_len, bpe->piece_str[ra], (size_t)ra_len);
np[la_len + ra_len] = '\0';
bpe->piece_str[vocab_count] = np;
bpe->piece_len[vocab_count] = la_len + ra_len;
bpe->piece_score[vocab_count] = -(float)vocab_count;
bpe->piece_type[vocab_count] = 0;
int new_id = vocab_count;
vocab_count++;
for (int s = 0; s < ns; s++) {
if (!syms[s].active) continue;
int nxt = syms[s].next;
if (nxt == -1 || !syms[nxt].active || syms[s].id != la || syms[nxt].id != ra) continue;
int nn = syms[nxt].next;
int prv = syms[s].prev;
if (prv != -1 && syms[prv].active) bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, syms[prv].id, la, -1);
bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, la, ra, -1);
if (nn != -1 && syms[nn].active) bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, ra, syms[nn].id, -1);
syms[s].id = new_id;
syms[s].next = nn;
syms[nxt].active = 0;
if (nn != -1) syms[nn].prev = s;
if (prv != -1 && syms[prv].active && !bpe_pair_in_vocab(bpe, syms[prv].id, new_id)) bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, syms[prv].id, new_id, 1);
if (nn != -1 && syms[nn].active && !bpe_pair_in_vocab(bpe, new_id, syms[nn].id)) bpe_map_set(&heap, &heap_n, &heap_cap, bigrams, bigram_mask, new_id, syms[nn].id, 1);
}
merges_done++;
}
free(bigrams);
free(heap);
free(syms);
free(buf);
bpe->vocab_size = vocab_count;
}
static int *bpe_encode(const BPEVocab *bpe, const char *text, size_t text_len, int *out_len) {
int tl = (int)text_len;
char *buf = (char *)malloc((size_t)tl * 3 + 4);
int bn = 0;
buf[bn++] = (char)0xE2; buf[bn++] = (char)0x96; buf[bn++] = (char)0x81;
for (int i = 0; i < tl; ++i) {
if (text[i] == ' ') { buf[bn++] = (char)0xE2; buf[bn++] = (char)0x96; buf[bn++] = (char)0x81; }
else buf[bn++] = text[i];
}
SymTok *sym = (SymTok *)malloc(sizeof(SymTok) * (size_t)(bn + 1));
int ns = 0;
for (int i = 0; i < bn; ) {
int l = utf8_len((unsigned char)buf[i]); if (i + l > bn) l = bn - i;
sym[ns].start = i; sym[ns].end = i + l;
sym[ns].prev = ns - 1; sym[ns].next = ns + 1; sym[ns].active = 1;
ns++; i += l;
}
if (ns) sym[ns - 1].next = -1;
BPEHeap heap = {0};
BPECand c;
for (int i = 0; i + 1 < ns; ++i) {
int id = bpe_lookup(bpe, buf + sym[i].start, sym[i + 1].end - sym[i].start);
if (id >= 0) { c.score = bpe->piece_score[id]; c.left = i; c.right = i + 1; bpe_heap_push(&heap, c); }
}
while (bpe_heap_pop(&heap, &c)) {
SymTok *L = &sym[c.left];
if (!L->active || L->next != c.right || !sym[c.right].active) continue;
SymTok *R = &sym[c.right];
int id = bpe_lookup(bpe, buf + L->start, R->end - L->start);
if (id < 0 || bpe->piece_score[id] != c.score) continue;
L->end = R->end; L->next = R->next;
if (R->next != -1) sym[R->next].prev = c.left;
R->active = 0;
if (L->prev != -1) {
int p = L->prev;
int pid = bpe_lookup(bpe, buf + sym[p].start, L->end - sym[p].start);
if (pid >= 0) { BPECand n2 = { bpe->piece_score[pid], p, c.left }; bpe_heap_push(&heap, n2); }
}
if (L->next != -1) {
int nid = bpe_lookup(bpe, buf + L->start, sym[L->next].end - L->start);
if (nid >= 0) { BPECand n2 = { bpe->piece_score[nid], c.left, L->next }; bpe_heap_push(&heap, n2); }
}
}
int *tokens = (int *)malloc((size_t)(ns + 1) * sizeof(int));
int no = 0;
for (int i = 0; i != -1 && i < ns; i = sym[i].next) {
if (!sym[i].active) continue;
int len = sym[i].end - sym[i].start;
int id = bpe_lookup(bpe, buf + sym[i].start, len);
if (id >= 0) { tokens[no++] = id; }
else {
for (int b = sym[i].start; b < sym[i].end; ++b) {
int bt = bpe->byte_tok[(unsigned char)buf[b]];
if (bt < 0) bt = bpe->unk_id;
tokens[no++] = bt;
}
}
}
free(heap.a); free(sym); free(buf);
*out_len = no;
return tokens;
}
static char *bpe_decode(const BPEVocab *bpe, const int *ids, int n) {
size_t cap = 256, len = 0;
char *raw = (char *)malloc(cap);
for (int i = 0; i < n; ++i) {
int id = ids[i];
if (id < 0 || id >= bpe->vocab_size) continue;
if (bpe->piece_type[id] == 3) {
const char *p = bpe->piece_str[id];
if (bpe->piece_len[id] >= 6) {
int h1 = p[3], h2 = p[4];
int hi = (h1 <= '9') ? h1 - '0' : (h1 | 32) - 'a' + 10;
int lo2 = (h2 <= '9') ? h2 - '0' : (h2 | 32) - 'a' + 10;
if (len + 1 > cap) { cap *= 2; raw = (char *)realloc(raw, cap); }
raw[len++] = (char)((hi << 4) | lo2);
}
} else if (bpe->piece_type[id] == 2) {
continue;
} else if (bpe->piece_type[id] == 1) {
if (len + 1 > cap) { cap *= 2; raw = (char *)realloc(raw, cap); }
raw[len++] = '?';
} else {
int pl = bpe->piece_len[id]; const char *p = bpe->piece_str[id];
if (len + (size_t)pl > cap) { while (len + (size_t)pl > cap) cap *= 2; raw = (char *)realloc(raw, cap); }
memcpy(raw + len, p, (size_t)pl); len += (size_t)pl;
}
}
char *out = (char *)malloc(len + 2);
size_t o = 0;
for (size_t i = 0; i < len; ) {
if (i + 2 < len && (unsigned char)raw[i] == 0xE2 && (unsigned char)raw[i+1] == 0x96 && (unsigned char)raw[i+2] == 0x81) {
out[o++] = ' '; i += 3;
} else out[o++] = raw[i++];
}
out[o] = 0;
free(raw);
if (out[0] == ' ') memmove(out, out + 1, o);
return out;
}
static void bpe_decode_token(const BPEVocab *bpe, int token, char *out, int out_size) {
int ids[1] = { token };
char *decoded = bpe_decode(bpe, ids, 1);
strncpy(out, decoded, (size_t)(out_size - 1));
out[out_size - 1] = '\0';
free(decoded);
}
static void free_bpe_vocab(BPEVocab *bpe) {
if (bpe->piece_str) {
for (int i = 0; i < bpe->vocab_size; i++) {
free(bpe->piece_str[i]);
}
free(bpe->piece_str);
}
free(bpe->piece_len);
free(bpe->piece_score);
free(bpe->piece_type);
free(bpe->hash_id);
memset(bpe, 0, sizeof(BPEVocab));
}
typedef struct {
char chars[MAX_VOCAB_SIZE];
int char_to_idx[256];
int size;
} Vocabulary;
#define MAX_WORD_LEN 64
#define MAX_WORDS 32768
#define WORD_HASH_SIZE 65536
typedef struct {
char **words;
int *hash_table;
int *hash_keys;
int size;
int capacity;
int unk_idx;
int pad_idx;
int space_idx;
int newline_idx;
} WordVocabulary;
typedef enum {
TOKENIZER_CHAR,
TOKENIZER_WORD,
TOKENIZER_AUTO,
TOKENIZER_BPE
} TokenizerType;
typedef struct {
TokenizerType type;
Vocabulary char_vocab;
WordVocabulary word_vocab;
BPEVocab bpe_vocab;
} Tokenizer;
static void init_word_vocabulary(WordVocabulary *wv);
static void free_word_vocabulary(WordVocabulary *wv);
static int word_vocab_find(const WordVocabulary *wv, const char *word);
static int word_vocab_add(WordVocabulary *wv, const char *word);
static void build_word_vocabulary(WordVocabulary *wv, const char *text, size_t len);
static int *tokenize_words(const char *text, size_t len, const WordVocabulary *wv, int *out_len);
static const char *decode_word_token(int token, const WordVocabulary *wv);
static void build_vocabulary(Vocabulary *vocab, const char *text, size_t len);
static int *tokenize(const char *text, size_t len, const Vocabulary *vocab, int *out_len);
static void init_tokenizer(Tokenizer *tok, TokenizerType type);
static void free_tokenizer(Tokenizer *tok);
static void build_tokenizer(Tokenizer *tok, const char *text, size_t len, TokenizerType requested_type);
static int tokenizer_vocab_size(const Tokenizer *tok);
static int *tokenizer_encode(const Tokenizer *tok, const char *text, size_t len, int *out_len);
static void tokenizer_decode_token(const Tokenizer *tok, int token, char *out, int out_size);
static void init_layer_params(LayerParams *lp, const lrnnConfig *cfg, int layer_idx) {
int n_embd = cfg->n_embd;
int ffn_h = ffn_hidden(cfg);
int lora_rank = cfg->decay_lora_rank;
float proj_scale = 0.02f / sqrtf((float)cfg->n_layer);
float ffn_scale = proj_scale;
int n_slots = cfg->n_mem_slots;
lp->mem_gate_write = tensor_alloc(n_embd, n_slots);
tensor_randn(&lp->mem_gate_write, 0.01f);
lp->mem_gate_read = tensor_alloc(n_embd, n_slots);
tensor_randn(&lp->mem_gate_read, 0.01f);
lp->ln1_weight = tensor_alloc_1d(n_embd);
tensor_fill(&lp->ln1_weight, 1.0f);
lp->ln1_bias = tensor_alloc_1d(n_embd);
lp->ln2_weight = tensor_alloc_1d(n_embd);
tensor_fill(&lp->ln2_weight, 1.0f);
lp->ln2_bias = tensor_alloc_1d(n_embd);
lp->time_shift_w1 = tensor_alloc_1d(n_embd);
tensor_rand_uniform(&lp->time_shift_w1, 0.3f, 0.7f);
lp->time_shift_w2 = tensor_alloc_1d(n_embd);
tensor_rand_uniform(&lp->time_shift_w2, 0.1f, 0.3f);
lp->time_shift_w4 = tensor_alloc_1d(n_embd);
tensor_rand_uniform(&lp->time_shift_w4, 0.0f, 0.2f);
lp->time_mix_r = tensor_alloc_1d(n_embd);
tensor_fill(&lp->time_mix_r, 0.5f);
lp->time_mix_k = tensor_alloc_1d(n_embd);
tensor_fill(&lp->time_mix_k, 0.5f);
lp->time_mix_v = tensor_alloc_1d(n_embd);
tensor_fill(&lp->time_mix_v, 0.5f);
lp->decay_lora_a = tensor_alloc(n_embd, lora_rank);
tensor_randn(&lp->decay_lora_a, 0.01f);
lp->decay_lora_b = tensor_alloc(lora_rank, n_embd);
tensor_randn(&lp->decay_lora_b, 0.01f);
lp->decay_base = tensor_alloc_1d(n_embd);
{
int hdim_val = n_embd / cfg->n_head;
for (int h = 0; h < cfg->n_head; h++) {
float base_val = 1.5f - 0.1f * layer_idx - 0.2f * h;
for (int d = 0; d < hdim_val; d++) {
lp->decay_base.data[h * hdim_val + d] = base_val;
}
}
}
lp->time_first = tensor_alloc_1d(n_embd);
{
int hdim_val = n_embd / cfg->n_head;
for (int h = 0; h < cfg->n_head; h++) {
float tf_val = -3.0f + layer_idx * 0.3f + h * 0.5f;
for (int d = 0; d < hdim_val; d++) {
lp->time_first.data[h * hdim_val + d] = tf_val;
}
}
}
lp->Wr = tensor_alloc(n_embd, n_embd);
tensor_randn(&lp->Wr, proj_scale);
lp->Wk = tensor_alloc(n_embd, n_embd);
tensor_randn(&lp->Wk, proj_scale);
lp->Wv = tensor_alloc(n_embd, n_embd);
tensor_randn(&lp->Wv, proj_scale);
lp->Wo = tensor_alloc(n_embd, n_embd);
tensor_randn(&lp->Wo, proj_scale);
lp->channel_mix = tensor_alloc_1d(n_embd);
tensor_fill(&lp->channel_mix, 0.5f);
lp->ffn_gate_up = tensor_alloc(n_embd, 2 * ffn_h);
for (int k = 0; k < n_embd; k++) {
for (int j = 0; j < ffn_h; j++) lp->ffn_gate_up.data[(size_t)k * 2 * ffn_h + j] = randn() * ffn_scale;
}
for (int k = 0; k < n_embd; k++) {
for (int j = 0; j < ffn_h; j++) lp->ffn_gate_up.data[(size_t)k * 2 * ffn_h + ffn_h + j] = randn() * ffn_scale;
}
lp->ffn_down = tensor_alloc(ffn_h, n_embd);
tensor_randn(&lp->ffn_down, ffn_scale);
lp->alibi_slopes = tensor_alloc_1d(cfg->n_head);
compute_alibi_slopes(lp->alibi_slopes.data, cfg->n_head);
}
static void free_layer_params(LayerParams *lp) {
tensor_free(&lp->ln1_weight); tensor_free(&lp->ln1_bias);
tensor_free(&lp->ln2_weight); tensor_free(&lp->ln2_bias);
tensor_free(&lp->time_shift_w1); tensor_free(&lp->time_shift_w2); tensor_free(&lp->time_shift_w4);
tensor_free(&lp->time_mix_r); tensor_free(&lp->time_mix_k); tensor_free(&lp->time_mix_v);
tensor_free(&lp->decay_lora_a); tensor_free(&lp->decay_lora_b);
tensor_free(&lp->decay_base); tensor_free(&lp->time_first);
tensor_free(&lp->Wr); tensor_free(&lp->Wk); tensor_free(&lp->Wv); tensor_free(&lp->Wo);
tensor_free(&lp->channel_mix);
tensor_free(&lp->ffn_gate_up); tensor_free(&lp->ffn_down);
tensor_free(&lp->alibi_slopes);
tensor_free(&lp->mem_gate_write); tensor_free(&lp->mem_gate_read);
}
static void init_model_params(ModelParams *mp, const lrnnConfig *cfg) {
int n_embd = cfg->n_embd;
int vocab_size = cfg->vocab_size;
if (n_embd % cfg->n_head != 0) {
fprintf(stderr, "Error: n_embd (%d) must be divisible by n_head (%d)\n", n_embd, cfg->n_head);
exit(1);
}
mp->n_layers = cfg->n_layer;
mp->emb = tensor_alloc(vocab_size, n_embd);
tensor_randn(&mp->emb, 0.02f);
mp->ln0_weight = tensor_alloc_1d(n_embd);
tensor_fill(&mp->ln0_weight, 1.0f);
mp->ln0_bias = tensor_alloc_1d(n_embd);
mp->layers = (LayerParams *)calloc((size_t)cfg->n_layer, sizeof(LayerParams));
if (!mp->layers) { fprintf(stderr, "Error: failed to allocate layers\n"); exit(1); }
for (int i = 0; i < cfg->n_layer; i++) {
init_layer_params(&mp->layers[i], cfg, i);
}
mp->ln_out_weight = tensor_alloc_1d(n_embd);
tensor_fill(&mp->ln_out_weight, 1.0f);
mp->ln_out_bias = tensor_alloc_1d(n_embd);
mp->head = tensor_alloc(n_embd, vocab_size);
tensor_randn(&mp->head, 0.02f);
}
static void free_model_params(ModelParams *mp) {
tensor_free(&mp->emb);
tensor_free(&mp->ln0_weight); tensor_free(&mp->ln0_bias);
for (int i = 0; i < mp->n_layers; i++) free_layer_params(&mp->layers[i]);
free(mp->layers); mp->layers = NULL;
tensor_free(&mp->ln_out_weight); tensor_free(&mp->ln_out_bias);
tensor_free(&mp->head);
}
static void init_layer_state(LayerState *ls, int n_embd, int n_mem_slots) {
ls->x_prev_1 = tensor_alloc_1d(n_embd);
ls->x_prev_2 = tensor_alloc_1d(n_embd);
ls->x_prev_3 = tensor_alloc_1d(n_embd);
ls->x_prev_4 = tensor_alloc_1d(n_embd);
ls->wkv_num = tensor_alloc(n_mem_slots, n_embd);
ls->wkv_den = tensor_alloc(n_mem_slots, n_embd);
ls->ffn_prev = tensor_alloc_1d(n_embd);
}
static void free_layer_state(LayerState *ls) {
tensor_free(&ls->x_prev_1); tensor_free(&ls->x_prev_2);
tensor_free(&ls->x_prev_3); tensor_free(&ls->x_prev_4);
tensor_free(&ls->wkv_num); tensor_free(&ls->wkv_den);
tensor_free(&ls->ffn_prev);
}
static void init_model_state(ModelState *ms, const lrnnConfig *cfg) {
ms->n_layers = cfg->n_layer;
ms->layers = (LayerState *)calloc((size_t)cfg->n_layer, sizeof(LayerState));
if (!ms->layers) { fprintf(stderr, "Error: failed to allocate state\n"); exit(1); }
for (int i = 0; i < cfg->n_layer; i++) {
init_layer_state(&ms->layers[i], cfg->n_embd, cfg->n_mem_slots);
}
int n_embd = cfg->n_embd, ffn_h = ffn_hidden(cfg), n_slots = cfg->n_mem_slots;
ms->scr_n_embd = n_embd; ms->scr_ffn_h = ffn_h; ms->scr_n_slots = n_slots;
size_t total = (size_t)20 * n_embd + (size_t)3 * ffn_h + (size_t)4 * n_slots;
ms->scr = (float *)xaligned(total * sizeof(float));
}
static void free_model_state(ModelState *ms) {
for (int i = 0; i < ms->n_layers; i++) free_layer_state(&ms->layers[i]);
free(ms->layers); ms->layers = NULL;
free(ms->scr); ms->scr = NULL;
}
static inline int head_dim(const lrnnConfig *cfg) {
return cfg->n_embd / cfg->n_head;
}
static void forward_single(float *logits, int token, const ModelParams *mp,
ModelState *state, const lrnnConfig *cfg) {
int n_embd = cfg->n_embd;
int n_head = cfg->n_head;
int hdim = n_embd / n_head;
int ffn_h = ffn_hidden(cfg);
int n_slots = cfg->n_mem_slots;
float *sc = state->scr;
float *x = sc;
float *x_norm = sc + 1 * n_embd;
float *x_shifted = sc + 2 * n_embd;
float *xr = sc + 3 * n_embd;
float *xk = sc + 4 * n_embd;
float *xv = sc + 5 * n_embd;
float *r = sc + 6 * n_embd;
float *k = sc + 7 * n_embd;
float *v = sc + 8 * n_embd;
float *decay_delta = sc + 9 * n_embd;
float *decay = sc + 10 * n_embd;
float *k_exp = sc + 11 * n_embd;
float *time_first_val = sc + 12 * n_embd;
float *wkv = sc + 13 * n_embd;
float *tm_out = sc + 14 * n_embd;
float *xm = sc + 15 * n_embd;
float *cm_out = sc + 16 * n_embd;
float *lora_tmp = sc + 17 * n_embd;
float *w1_sig = sc + 18 * n_embd;
float *w2_sig = sc + 19 * n_embd;
float *w4_sig = sc + 20 * n_embd;
float *gate = sc + 20 * n_embd;
float *up = gate + ffn_h;
float *hidden = up + ffn_h;
float *write_logits = hidden + ffn_h;
float *write_gates = write_logits + n_slots;
float *read_logits = write_gates + n_slots;
float *read_gates = read_logits + n_slots;
memcpy(x, mp->emb.data + token * n_embd, (size_t)n_embd * sizeof(float));
layer_norm(x, x, mp->ln0_weight.data, mp->ln0_bias.data, n_embd);
for (int layer_idx = 0; layer_idx < mp->n_layers; layer_idx++) {
const LayerParams *lp = &mp->layers[layer_idx];
LayerState *ls = &state->layers[layer_idx];
layer_norm(x_norm, x, lp->ln1_weight.data, lp->ln1_bias.data, n_embd);
sigmoid_vec(w1_sig, lp->time_shift_w1.data, n_embd);
sigmoid_vec(w2_sig, lp->time_shift_w2.data, n_embd);
sigmoid_vec(w4_sig, lp->time_shift_w4.data, n_embd);
for (int i = 0; i < n_embd; i++) {
float w_sum = w1_sig[i] + w2_sig[i] + w4_sig[i] + EPSILON;
float nw1 = w1_sig[i] / w_sum;
float nw2 = w2_sig[i] / w_sum;
float nw4 = w4_sig[i] / w_sum;
x_shifted[i] = nw1 * ls->x_prev_1.data[i] +
nw2 * ls->x_prev_2.data[i] +
nw4 * ls->x_prev_4.data[i];
}
for (int i = 0; i < n_embd; i++) {
float mr = sigmoid_f(lp->time_mix_r.data[i]);
float mk = sigmoid_f(lp->time_mix_k.data[i]);
float mv = sigmoid_f(lp->time_mix_v.data[i]);
xr[i] = x_norm[i] * mr + x_shifted[i] * (1.0f - mr);
xk[i] = x_norm[i] * mk + x_shifted[i] * (1.0f - mk);
xv[i] = x_norm[i] * mv + x_shifted[i] * (1.0f - mv);
}
matvec(r, xr, &lp->Wr);
matvec(k, xk, &lp->Wk);
matvec(v, xv, &lp->Wv);
matvec(lora_tmp, x_norm, &lp->decay_lora_a);
matvec(decay_delta, lora_tmp, &lp->decay_lora_b);
for (int i = 0; i < n_embd; i++) {
decay[i] = sigmoid_f(lp->decay_base.data[i] + decay_delta[i]);
}
sigmoid_vec(r, r, n_embd);
exp_vec(time_first_val, lp->time_first.data, n_embd);
exp_vec(k_exp, k, n_embd);
{
matvec(write_logits, x_norm, &lp->mem_gate_write);
softmax_vec(write_gates, write_logits, n_slots);
matvec(read_logits, x_norm, &lp->mem_gate_read);
softmax_vec(read_gates, read_logits, n_slots);
for (int h = 0; h < n_head; h++) {
int base = h * hdim;
float alibi_decay_h = expf(-lp->alibi_slopes.data[h]);
for (int d = 0; d < hdim; d++) {
int i = base + d;
float kv = k_exp[i] * v[i];
float read_num = 0.0f, read_den = 0.0f;
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
read_num += read_gates[s] * ls->wkv_num.data[si];
read_den += read_gates[s] * ls->wkv_den.data[si];
}
float num = read_num + time_first_val[i] * kv;
float den = read_den + time_first_val[i] * k_exp[i] + EPSILON;
wkv[i] = num / den;
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
float wg = write_gates[s];
float combined = decay[i] * alibi_decay_h;
ls->wkv_num.data[si] = combined * ls->wkv_num.data[si] + wg * kv;
ls->wkv_den.data[si] = combined * ls->wkv_den.data[si] + wg * k_exp[i];
}
}
}
}
vec_mul(wkv, r, wkv, n_embd);
matvec(tm_out, wkv, &lp->Wo);
vec_add(x, x, tm_out, n_embd);
tensor_copy(&ls->x_prev_4, &ls->x_prev_3);
tensor_copy(&ls->x_prev_3, &ls->x_prev_2);
tensor_copy(&ls->x_prev_2, &ls->x_prev_1);
memcpy(ls->x_prev_1.data, x_norm, (size_t)n_embd * sizeof(float));
layer_norm(x_norm, x, lp->ln2_weight.data, lp->ln2_bias.data, n_embd);
for (int i = 0; i < n_embd; i++) {
float mix = sigmoid_f(lp->channel_mix.data[i]);
xm[i] = x_norm[i] * mix + ls->ffn_prev.data[i] * (1.0f - mix);
}
{
hp_gemv_block(xm, lp->ffn_gate_up.data, gate, 2 * ffn_h, n_embd, 0, ffn_h);
hp_gemv_block(xm, lp->ffn_gate_up.data, up - ffn_h, 2 * ffn_h, n_embd, ffn_h, 2 * ffn_h);
}
for (int i = 0; i < ffn_h; i++) {
hidden[i] = silu_f(gate[i]) * up[i];
}
matvec(cm_out, hidden, &lp->ffn_down);
vec_add(x, x, cm_out, n_embd);
memcpy(ls->ffn_prev.data, x_norm, (size_t)n_embd * sizeof(float));
}
layer_norm(x, x, mp->ln_out_weight.data, mp->ln_out_bias.data, n_embd);
matvec(logits, x, &mp->head);
}
static float cross_entropy_loss(const Tensor *logits, const int *targets, int n) {
int vocab_size = logits->cols;
float *probs = (float *)malloc((size_t)vocab_size * sizeof(float));
float total_loss = 0.0f;
for (int t = 0; t < n; t++) {
softmax_vec(probs, logits->data + t * vocab_size, vocab_size);
int target = targets[t];
float p = probs[target];
if (p < EPSILON) p = EPSILON;
total_loss -= logf(p);
}
free(probs);
return total_loss / (float)n;
}
static uint16_t f32_to_f16(float f) {
union { float f; uint32_t u; } c = { f };
uint32_t x = c.u;
uint16_t s = (uint16_t)((x >> 16) & 0x8000u);
int32_t e = (int32_t)((x >> 23) & 0xFF) - 127 + 15;
uint32_t m = x & 0x7FFFFFu;
if (e >= 31) return (uint16_t)(s | 0x7C00u);
if (e <= 0) {
if (e < -10) return s;
m |= 0x800000u;
uint32_t shift = (uint32_t)(14 - e);
return (uint16_t)(s | (m >> shift));
}
return (uint16_t)(s | ((uint32_t)e << 10) | (m >> 13));
}
static float f16_to_f32(uint16_t h) {
uint32_t s = (uint32_t)(h & 0x8000u) << 16;
uint32_t e = (h >> 10) & 0x1Fu, m = h & 0x3FFu;
uint32_t u;
if (e == 0) {
if (m == 0) u = s;
else {
int shift = 0;
while (!(m & 0x400u) && shift < 10) { m <<= 1; shift++; }
u = s | ((uint32_t)(113 - shift) << 23) | ((m & 0x3FFu) << 13);
}
} else if (e == 31) u = s | 0x7F800000u | (m << 13);
else u = s | ((e + 112) << 23) | (m << 13);
union { float f; uint32_t u; } c = { .u = u };
return c.f;
}
static void write_tensor_fp16(FILE *f, const Tensor *t) {
fwrite(&t->rows, sizeof(int), 1, f);
fwrite(&t->cols, sizeof(int), 1, f);
uint16_t *h = (uint16_t *)malloc((size_t)t->size * sizeof(uint16_t));
if ((size_t)t->size >= 8) {
long i = 0, n8 = (long)t->size - 7;
for (; i < n8; i += 8) {
__m128i hv = _mm256_cvtps_ph(_mm256_loadu_ps(t->data + i), 0);
_mm_storeu_si128((__m128i *)(h + i), hv);
}
for (; i < t->size; i++) h[i] = f32_to_f16(t->data[i]);
} else {
for (int i = 0; i < t->size; i++) h[i] = f32_to_f16(t->data[i]);
}
fwrite(h, sizeof(uint16_t), (size_t)t->size, f);
free(h);
}
static void read_tensor_fp16(FILE *f, Tensor *t) {
int rows, cols;
if (fread(&rows, sizeof(int), 1, f) != 1) return;
if (fread(&cols, sizeof(int), 1, f) != 1) return;
*t = tensor_alloc(rows, cols);
uint16_t *h = (uint16_t *)malloc((size_t)t->size * sizeof(uint16_t));
if (fread(h, sizeof(uint16_t), (size_t)t->size, f) != (size_t)t->size) {
fprintf(stderr, "Warning: incomplete tensor read\n");
}
if ((size_t)t->size >= 8) {
long i = 0, n8 = (long)t->size - 7;
for (; i < n8; i += 8) {
__m256 v = _mm256_cvtph_ps(_mm_loadu_si128((const __m128i *)(h + i)));
_mm256_storeu_ps(t->data + i, v);
}
for (; i < t->size; i++) t->data[i] = f16_to_f32(h[i]);
} else {
for (int i = 0; i < t->size; i++) t->data[i] = f16_to_f32(h[i]);
}
free(h);
}
static void read_tensor(FILE *f, Tensor *t) {
int rows, cols;
if (fread(&rows, sizeof(int), 1, f) != 1) return;
if (fread(&cols, sizeof(int), 1, f) != 1) return;
*t = tensor_alloc(rows, cols);
if (fread(t->data, sizeof(float), (size_t)t->size, f) != (size_t)t->size) {
fprintf(stderr, "Warning: incomplete tensor read\n");
}
}
#define read_tensor_any(f, t) (fp16_storage ? read_tensor_fp16((f), (t)) : read_tensor((f), (t)))
static int save_model(const char *path, const ModelParams *mp,
const lrnnConfig *cfg, const Tokenizer *tok) {
FILE *f = fopen(path, "wb");
if (!f) { fprintf(stderr, "Error: cannot open %s for writing\n", path); return -1; }
const char magic[] = "lrnnC04";
fwrite(magic, 1, 8, f);
fwrite(cfg, sizeof(lrnnConfig), 1, f);
int tok_type = (int)tok->type;
fwrite(&tok_type, sizeof(int), 1, f);
if (tok->type == TOKENIZER_CHAR) {
fwrite(&tok->char_vocab.size, sizeof(int), 1, f);
fwrite(tok->char_vocab.chars, sizeof(char), (size_t)tok->char_vocab.size, f);
fwrite(tok->char_vocab.char_to_idx, sizeof(int), 256, f);
} else if (tok->type == TOKENIZER_WORD) {
fwrite(&tok->word_vocab.size, sizeof(int), 1, f);
fwrite(&tok->word_vocab.unk_idx, sizeof(int), 1, f);
fwrite(&tok->word_vocab.pad_idx, sizeof(int), 1, f);
fwrite(&tok->word_vocab.space_idx, sizeof(int), 1, f);
fwrite(&tok->word_vocab.newline_idx, sizeof(int), 1, f);
for (int i = 0; i < tok->word_vocab.size; i++) {
int len = (int)strlen(tok->word_vocab.words[i]);
fwrite(&len, sizeof(int), 1, f);
fwrite(tok->word_vocab.words[i], sizeof(char), (size_t)len, f);
}
} else if (tok->type == TOKENIZER_BPE) {
const BPEVocab *bpe = &tok->bpe_vocab;
fwrite(&bpe->vocab_size, sizeof(int), 1, f);
fwrite(&bpe->unk_id, sizeof(int), 1, f);
fwrite(&bpe->pad_id, sizeof(int), 1, f);
for (int i = 0; i < bpe->vocab_size; i++) {
fwrite(&bpe->piece_type[i], sizeof(int8_t), 1, f);
fwrite(&bpe->piece_len[i], sizeof(int), 1, f);
fwrite(bpe->piece_str[i], sizeof(char), (size_t)bpe->piece_len[i], f);
fwrite(&bpe->piece_score[i], sizeof(float), 1, f);
}
}
write_tensor_fp16(f, &mp->emb);
write_tensor_fp16(f, &mp->ln0_weight); write_tensor_fp16(f, &mp->ln0_bias);
fwrite(&mp->n_layers, sizeof(int), 1, f);
for (int i = 0; i < mp->n_layers; i++) {
const LayerParams *lp = &mp->layers[i];
write_tensor_fp16(f, &lp->ln1_weight); write_tensor_fp16(f, &lp->ln1_bias);
write_tensor_fp16(f, &lp->ln2_weight); write_tensor_fp16(f, &lp->ln2_bias);
write_tensor_fp16(f, &lp->time_shift_w1); write_tensor_fp16(f, &lp->time_shift_w2); write_tensor_fp16(f, &lp->time_shift_w4);
write_tensor_fp16(f, &lp->time_mix_r); write_tensor_fp16(f, &lp->time_mix_k); write_tensor_fp16(f, &lp->time_mix_v);
write_tensor_fp16(f, &lp->decay_lora_a); write_tensor_fp16(f, &lp->decay_lora_b);
write_tensor_fp16(f, &lp->decay_base); write_tensor_fp16(f, &lp->time_first);
write_tensor_fp16(f, &lp->Wr); write_tensor_fp16(f, &lp->Wk); write_tensor_fp16(f, &lp->Wv); write_tensor_fp16(f, &lp->Wo);
write_tensor_fp16(f, &lp->channel_mix);
{
int fh = lp->ffn_gate_up.cols / 2;
Tensor gate = tensor_alloc(lp->ffn_gate_up.rows, fh), up = tensor_alloc(lp->ffn_gate_up.rows, fh);
for (int k = 0; k < lp->ffn_gate_up.rows; k++) {
memcpy(gate.data + (size_t)k * fh, lp->ffn_gate_up.data + (size_t)k * 2 * fh, (size_t)fh * sizeof(float));
memcpy(up.data + (size_t)k * fh, lp->ffn_gate_up.data + (size_t)k * 2 * fh + fh, (size_t)fh * sizeof(float));
}
write_tensor_fp16(f, &gate);
write_tensor_fp16(f, &up);
tensor_free(&gate); tensor_free(&up);
}
write_tensor_fp16(f, &lp->ffn_down);
write_tensor_fp16(f, &lp->mem_gate_write); write_tensor_fp16(f, &lp->mem_gate_read);
}
write_tensor_fp16(f, &mp->ln_out_weight); write_tensor_fp16(f, &mp->ln_out_bias);
write_tensor_fp16(f, &mp->head);
fclose(f);
return 0;
}
static int load_model(const char *path, ModelParams *mp,
lrnnConfig *cfg, Tokenizer *tok) {
FILE *f = fopen(path, "rb");
if (!f) { fprintf(stderr, "Error: cannot open %s for reading\n", path); return -1; }
char magic[8];
if (fread(magic, 1, 8, f) != 8) { fclose(f); return -1; }
bool is_v1 = (strncmp(magic, "lrnnC01", 7) == 0);
bool is_v2 = (strncmp(magic, "lrnnC02", 7) == 0);
bool is_v3 = (strncmp(magic, "lrnnC03", 7) == 0);
bool is_v4 = (strncmp(magic, "lrnnC04", 7) == 0);
bool fp16_storage = is_v4;
if (!is_v1 && !is_v2 && !is_v3 && !is_v4) {
fprintf(stderr, "Error: invalid model file format\n");
fclose(f); return -1;
}
if (fread(cfg, sizeof(lrnnConfig), 1, f) != 1) { fclose(f); return -1; }
memset(tok, 0, sizeof(Tokenizer));
if (is_v1) {
tok->type = TOKENIZER_CHAR;
if (fread(&tok->char_vocab.size, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(tok->char_vocab.chars, sizeof(char), (size_t)tok->char_vocab.size, f) != (size_t)tok->char_vocab.size) { fclose(f); return -1; }
if (fread(tok->char_vocab.char_to_idx, sizeof(int), 256, f) != 256) { fclose(f); return -1; }
} else {
int tok_type;
if (fread(&tok_type, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
tok->type = (TokenizerType)tok_type;
if (tok->type == TOKENIZER_CHAR) {
if (fread(&tok->char_vocab.size, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(tok->char_vocab.chars, sizeof(char), (size_t)tok->char_vocab.size, f) != (size_t)tok->char_vocab.size) { fclose(f); return -1; }
if (fread(tok->char_vocab.char_to_idx, sizeof(int), 256, f) != 256) { fclose(f); return -1; }
} else if (tok->type == TOKENIZER_WORD) {
init_word_vocabulary(&tok->word_vocab);
int size;
if (fread(&size, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&tok->word_vocab.unk_idx, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&tok->word_vocab.pad_idx, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&tok->word_vocab.space_idx, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&tok->word_vocab.newline_idx, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
for (int i = 0; i < size; i++) {
int len;
if (fread(&len, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
char word[MAX_WORD_LEN];
if (len >= MAX_WORD_LEN) len = MAX_WORD_LEN - 1;
if (fread(word, sizeof(char), (size_t)len, f) != (size_t)len) { fclose(f); return -1; }
word[len] = '\0';
word_vocab_add(&tok->word_vocab, word);
}
} else if (tok->type == TOKENIZER_BPE) {
BPEVocab *bpe = &tok->bpe_vocab;
if (fread(&bpe->vocab_size, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&bpe->unk_id, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
if (fread(&bpe->pad_id, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
bpe->piece_str = (char **)calloc((size_t)bpe->vocab_size, sizeof(char *));
bpe->piece_len = (int *)calloc((size_t)bpe->vocab_size, sizeof(int));
bpe->piece_score = (float *)calloc((size_t)bpe->vocab_size, sizeof(float));
bpe->piece_type = (int8_t *)calloc((size_t)bpe->vocab_size, sizeof(int8_t));
for (int i = 0; i < 256; i++) bpe->byte_tok[i] = -1;
for (int i = 0; i < bpe->vocab_size; i++) {
if (fread(&bpe->piece_type[i], sizeof(int8_t), 1, f) != 1) { fclose(f); return -1; }
if (fread(&bpe->piece_len[i], sizeof(int), 1, f) != 1) { fclose(f); return -1; }
bpe->piece_str[i] = (char *)malloc((size_t)(bpe->piece_len[i] + 1));
if (fread(bpe->piece_str[i], sizeof(char), (size_t)bpe->piece_len[i], f) != (size_t)bpe->piece_len[i]) { fclose(f); return -1; }
bpe->piece_str[i][bpe->piece_len[i]] = '\0';
if (fread(&bpe->piece_score[i], sizeof(float), 1, f) != 1) { fclose(f); return -1; }
if (bpe->piece_type[i] == 3 && bpe->piece_len[i] == 6) {
const char *p = bpe->piece_str[i];
int h1 = p[3], h2 = p[4];
int hi = (h1 <= '9') ? h1 - '0' : (h1 | 32) - 'a' + 10;
int lo2 = (h2 <= '9') ? h2 - '0' : (h2 | 32) - 'a' + 10;
bpe->byte_tok[(hi << 4) | lo2] = i;
}
}
int hs = 1;
while (hs < bpe->vocab_size * 2) hs <<= 1;
bpe->hash_mask = hs - 1;
bpe->hash_id = (int *)malloc((size_t)hs * sizeof(int));
for (int i = 0; i < hs; i++) bpe->hash_id[i] = -1;
for (int i = 0; i < bpe->vocab_size; i++) {
if (bpe->piece_type[i] == 1 || bpe->piece_type[i] == 2) continue;
uint64_t h = fnv1a(bpe->piece_str[i], bpe->piece_len[i]) & (uint64_t)bpe->hash_mask;
while (bpe->hash_id[h] >= 0) h = (h + 1) & (uint64_t)bpe->hash_mask;
bpe->hash_id[h] = i;
}
}
}
read_tensor_any(f, &mp->emb);
read_tensor_any(f, &mp->ln0_weight); read_tensor_any(f, &mp->ln0_bias);
if (fread(&mp->n_layers, sizeof(int), 1, f) != 1) { fclose(f); return -1; }
mp->layers = (LayerParams *)calloc((size_t)mp->n_layers, sizeof(LayerParams));
for (int i = 0; i < mp->n_layers; i++) {
LayerParams *lp = &mp->layers[i];
read_tensor_any(f, &lp->ln1_weight); read_tensor_any(f, &lp->ln1_bias);
read_tensor_any(f, &lp->ln2_weight); read_tensor_any(f, &lp->ln2_bias);
read_tensor_any(f, &lp->time_shift_w1); read_tensor_any(f, &lp->time_shift_w2); read_tensor_any(f, &lp->time_shift_w4);
read_tensor_any(f, &lp->time_mix_r); read_tensor_any(f, &lp->time_mix_k); read_tensor_any(f, &lp->time_mix_v);
read_tensor_any(f, &lp->decay_lora_a); read_tensor_any(f, &lp->decay_lora_b);
read_tensor_any(f, &lp->decay_base); read_tensor_any(f, &lp->time_first);
read_tensor_any(f, &lp->Wr); read_tensor_any(f, &lp->Wk); read_tensor_any(f, &lp->Wv); read_tensor_any(f, &lp->Wo);
read_tensor_any(f, &lp->channel_mix);
{
Tensor gate = tensor_alloc(0, 0), up = tensor_alloc(0, 0);
read_tensor_any(f, &gate);
read_tensor_any(f, &up);
lp->ffn_gate_up = tensor_alloc(gate.rows, gate.cols + up.cols);
for (int k = 0; k < gate.rows; k++) {
memcpy(lp->ffn_gate_up.data + (size_t)k * (gate.cols + up.cols), gate.data + (size_t)k * gate.cols, (size_t)gate.cols * sizeof(float));
memcpy(lp->ffn_gate_up.data + (size_t)k * (gate.cols + up.cols) + gate.cols, up.data + (size_t)k * up.cols, (size_t)up.cols * sizeof(float));
}
tensor_free(&gate); tensor_free(&up);
}
read_tensor_any(f, &lp->ffn_down);
read_tensor_any(f, &lp->mem_gate_write); read_tensor_any(f, &lp->mem_gate_read);
}
for (int i = 0; i < mp->n_layers; i++) {
mp->layers[i].alibi_slopes = tensor_alloc_1d(cfg->n_head);
compute_alibi_slopes(mp->layers[i].alibi_slopes.data, cfg->n_head);
}
read_tensor_any(f, &mp->ln_out_weight); read_tensor_any(f, &mp->ln_out_bias);
read_tensor_any(f, &mp->head);
fclose(f);
return 0;
}
static void build_vocabulary(Vocabulary *vocab, const char *text, size_t len) {
bool seen[256] = {false};
vocab->size = 0;
for (size_t i = 0; i < len; i++) {
unsigned char c = (unsigned char)text[i];
if (!seen[c]) {
seen[c] = true;
vocab->chars[vocab->size] = (char)c;
vocab->char_to_idx[c] = vocab->size;
vocab->size++;
}
}
for (int i = 0; i < vocab->size - 1; i++) {
for (int j = i + 1; j < vocab->size; j++) {
if ((unsigned char)vocab->chars[i] > (unsigned char)vocab->chars[j]) {
char tmp = vocab->chars[i]; vocab->chars[i] = vocab->chars[j]; vocab->chars[j] = tmp;
}
}
}
for (int i = 0; i < 256; i++) vocab->char_to_idx[i] = 0;
for (int i = 0; i < vocab->size; i++) vocab->char_to_idx[(unsigned char)vocab->chars[i]] = i;
}
static unsigned int word_hash(const char *word) {
unsigned int hash = 5381;
while (*word) { hash = ((hash << 5) + hash) ^ (unsigned char)*word++; }
return hash;
}
static void init_word_vocabulary(WordVocabulary *wv) {
wv->capacity = MAX_WORDS;
wv->words = (char **)calloc((size_t)wv->capacity, sizeof(char *));
wv->hash_table = (int *)malloc(WORD_HASH_SIZE * sizeof(int));
wv->hash_keys = (int *)malloc(WORD_HASH_SIZE * sizeof(int));
for (int i = 0; i < WORD_HASH_SIZE; i++) { wv->hash_table[i] = -1; wv->hash_keys[i] = -1; }
wv->size = 0; wv->unk_idx = -1; wv->pad_idx = -1; wv->space_idx = -1; wv->newline_idx = -1;
}
static void free_word_vocabulary(WordVocabulary *wv) {
if (wv->words) { for (int i = 0; i < wv->size; i++) free(wv->words[i]); free(wv->words); wv->words = NULL; }
free(wv->hash_table); wv->hash_table = NULL;
free(wv->hash_keys); wv->hash_keys = NULL;
wv->size = 0;
}
static int word_vocab_find(const WordVocabulary *wv, const char *word) {
unsigned int hash = word_hash(word);
unsigned int idx = hash % WORD_HASH_SIZE;
for (int probe = 0; probe < 1000; probe++) {
unsigned int slot = (idx + (unsigned)probe) % WORD_HASH_SIZE;
if (wv->hash_table[slot] < 0) return -1;
if (wv->hash_keys[slot] == (int)hash) {
int word_idx = wv->hash_table[slot];
if (strcmp(wv->words[word_idx], word) == 0) return word_idx;
}
}
return -1;
}
static int word_vocab_add(WordVocabulary *wv, const char *word) {
int existing = word_vocab_find(wv, word);
if (existing >= 0) return existing;
if (wv->size >= wv->capacity - 1) { fprintf(stderr, "Warning: word vocabulary full\n"); return wv->unk_idx; }
int word_idx = wv->size;
wv->words[word_idx] = strdup(word);
wv->size++;
unsigned int hash = word_hash(word);
unsigned int idx = hash % WORD_HASH_SIZE;
for (int probe = 0; probe < 1000; probe++) {
unsigned int slot = (idx + (unsigned)probe) % WORD_HASH_SIZE;
if (wv->hash_table[slot] < 0) { wv->hash_table[slot] = word_idx; wv->hash_keys[slot] = (int)hash; break; }
}
return word_idx;
}
static inline int is_word_boundary(char c) {
return c == ' ' || c == '\n' || c == '\t' || c == '\r' ||
c == '.' || c == ',' || c == '!' || c == '?' ||
c == ':' || c == ';' || c == '"' || c == '\'' ||
c == '(' || c == ')' || c == '[' || c == ']' ||
c == '{' || c == '}' || c == '-' || c == '/' ||
c == '\\' || c == '@' || c == '#' || c == '$' ||
c == '%' || c == '&' || c == '*' || c == '+' ||
c == '=' || c == '<' || c == '>' || c == '|' ||
c == '~' || c == '`' || c == '^';
}
static void build_word_vocabulary(WordVocabulary *wv, const char *text, size_t len) {
init_word_vocabulary(wv);
wv->unk_idx = word_vocab_add(wv, "");
wv->pad_idx = word_vocab_add(wv, "");
wv->space_idx = word_vocab_add(wv, " ");
wv->newline_idx = word_vocab_add(wv, "\n");
word_vocab_add(wv, "."); word_vocab_add(wv, ","); word_vocab_add(wv, "!");
word_vocab_add(wv, "?"); word_vocab_add(wv, ":"); word_vocab_add(wv, ";");
word_vocab_add(wv, "\""); word_vocab_add(wv, "'"); word_vocab_add(wv, "(");
word_vocab_add(wv, ")"); word_vocab_add(wv, "-"); word_vocab_add(wv, "\t");
char word[MAX_WORD_LEN]; int word_len = 0;
for (size_t i = 0; i < len; i++) {
char c = text[i];
if (is_word_boundary(c)) {
if (word_len > 0) { word[word_len] = '\0'; word_vocab_add(wv, word); word_len = 0; }
if (c != ' ' && c != '\t' && c != '\r') { char punct[2] = {c, '\0'}; word_vocab_add(wv, punct); }
} else { if (word_len < MAX_WORD_LEN - 1) word[word_len++] = c; }
}
if (word_len > 0) { word[word_len] = '\0'; word_vocab_add(wv, word); }
}
static int *tokenize_words(const char *text, size_t len, const WordVocabulary *wv, int *out_len) {
int max_tokens = (int)(len / 2) + 100;
int *tokens = (int *)malloc((size_t)max_tokens * sizeof(int));
int n_tokens = 0;
char word[MAX_WORD_LEN]; int word_len = 0;
for (size_t i = 0; i < len; i++) {
char c = text[i];
if (is_word_boundary(c)) {
if (word_len > 0) { word[word_len] = '\0'; int idx = word_vocab_find(wv, word); tokens[n_tokens++] = (idx >= 0) ? idx : wv->unk_idx; word_len = 0; }
if (c == ' ') tokens[n_tokens++] = wv->space_idx;
else if (c == '\n') tokens[n_tokens++] = wv->newline_idx;
else if (c == '\t') tokens[n_tokens++] = wv->space_idx;
else if (c != '\r') { char punct[2] = {c, '\0'}; int idx = word_vocab_find(wv, punct); if (idx >= 0) tokens[n_tokens++] = idx; }
} else { if (word_len < MAX_WORD_LEN - 1) word[word_len++] = c; }
if (n_tokens >= max_tokens - 10) { max_tokens *= 2; tokens = (int *)realloc(tokens, (size_t)max_tokens * sizeof(int)); }
}
if (word_len > 0) { word[word_len] = '\0'; int idx = word_vocab_find(wv, word); tokens[n_tokens++] = (idx >= 0) ? idx : wv->unk_idx; }
*out_len = n_tokens;
return tokens;
}
static const char *decode_word_token(int token, const WordVocabulary *wv) {
if (token >= 0 && token < wv->size && wv->words[token]) return wv->words[token];
return "";
}
static void init_tokenizer(Tokenizer *tok, TokenizerType type) {
memset(tok, 0, sizeof(Tokenizer));
tok->type = type;
}
static void free_tokenizer(Tokenizer *tok) {
if (tok->type == TOKENIZER_WORD) free_word_vocabulary(&tok->word_vocab);
else if (tok->type == TOKENIZER_BPE) free_bpe_vocab(&tok->bpe_vocab);
}
static void build_tokenizer(Tokenizer *tok, const char *text, size_t len, TokenizerType requested_type) {
if (requested_type == TOKENIZER_AUTO) {
if (len < 5000) {
tok->type = TOKENIZER_CHAR;
printf(" Auto-selected: character tokenizer (corpus < 5KB)\n");
} else {
tok->type = TOKENIZER_BPE;
printf(" Auto-selected: BPE tokenizer (production, corpus >= 5KB)\n");
}
} else {
tok->type = requested_type;
}
if (tok->type == TOKENIZER_CHAR) {
build_vocabulary(&tok->char_vocab, text, len);
} else if (tok->type == TOKENIZER_WORD) {
build_word_vocabulary(&tok->word_vocab, text, len);
} else if (tok->type == TOKENIZER_BPE) {
int target_vs = 2048;
if (len > 100000) target_vs = 4096;
if (len > 1000000) target_vs = 8192;
bpe_build_vocab(&tok->bpe_vocab, text, len, target_vs);
}
}
static int tokenizer_vocab_size(const Tokenizer *tok) {
if (tok->type == TOKENIZER_CHAR) return tok->char_vocab.size;
else if (tok->type == TOKENIZER_WORD) return tok->word_vocab.size;
else return tok->bpe_vocab.vocab_size;
}
static int *tokenize(const char *text, size_t len, const Vocabulary *vocab, int *out_len) {
int *tokens = (int *)malloc(len * sizeof(int));
for (size_t i = 0; i < len; i++) tokens[i] = vocab->char_to_idx[(unsigned char)text[i]];
*out_len = (int)len;
return tokens;
}
static int *tokenizer_encode(const Tokenizer *tok, const char *text, size_t len, int *out_len) {
if (tok->type == TOKENIZER_CHAR) return tokenize(text, len, &tok->char_vocab, out_len);
else if (tok->type == TOKENIZER_WORD) return tokenize_words(text, len, &tok->word_vocab, out_len);
else return bpe_encode(&tok->bpe_vocab, text, len, out_len);
}
static void tokenizer_decode_token(const Tokenizer *tok, int token, char *out, int out_size) {
if (tok->type == TOKENIZER_CHAR) {
if (token >= 0 && token < tok->char_vocab.size) { out[0] = tok->char_vocab.chars[token]; out[1] = '\0'; }
else { out[0] = '?'; out[1] = '\0'; }
} else if (tok->type == TOKENIZER_WORD) {
const char *word = decode_word_token(token, &tok->word_vocab);
strncpy(out, word, (size_t)(out_size - 1)); out[out_size - 1] = '\0';
} else {
bpe_decode_token(&tok->bpe_vocab, token, out, out_size);
}
}
static lrnnConfig config_for_corpus(long corpus_bytes, TokenizerType tok_type, int vocab_size) {
lrnnConfig cfg = default_config();
cfg.vocab_size = vocab_size;
float efficiency = (tok_type == TOKENIZER_WORD || tok_type == TOKENIZER_BPE) ? 5.0f : 1.0f;
long effective_size = (long)((float)corpus_bytes / efficiency);
if (effective_size < 5000) {
cfg.n_layer = 2; cfg.n_embd = 64; cfg.ctx_len = 64;
cfg.decay_lora_rank = 4; cfg.ffn_multiplier = 1.5f; cfg.n_mem_slots = 2;
} else if (effective_size < 50000) {
cfg.n_layer = 4; cfg.n_embd = 128; cfg.ctx_len = 128;
cfg.decay_lora_rank = 8; cfg.ffn_multiplier = 2.0f; cfg.n_mem_slots = 4;
} else if (effective_size < 500000) {
cfg.n_layer = 6; cfg.n_embd = 256; cfg.ctx_len = 256;
cfg.decay_lora_rank = 16; cfg.ffn_multiplier = 2.5f; cfg.n_mem_slots = 4;
} else {
cfg.n_layer = 8; cfg.n_embd = 384; cfg.ctx_len = 512;
cfg.decay_lora_rank = 32; cfg.ffn_multiplier = 3.0f; cfg.n_mem_slots = 8;
}
cfg.n_head = cfg.n_embd / 32;
if (cfg.n_head < 2) cfg.n_head = 2;
return cfg;
}
typedef struct {
Tensor ln1_weight, ln1_bias, ln2_weight, ln2_bias;
Tensor time_shift_w1, time_shift_w2, time_shift_w4;
Tensor time_mix_r, time_mix_k, time_mix_v;
Tensor decay_lora_a, decay_lora_b, decay_base, time_first;
Tensor Wr, Wk, Wv, Wo;
Tensor channel_mix;
Tensor ffn_gate_up, ffn_down;
Tensor mem_gate_write, mem_gate_read;
} LayerGrads;
typedef struct {
Tensor emb;
Tensor ln0_weight, ln0_bias;
LayerGrads *layers;
Tensor ln_out_weight, ln_out_bias;
Tensor head;
int n_layers;
} ModelGrads;
typedef struct {
Tensor x_in;
Tensor x_ln1, x_shifted;
Tensor shift_w1_sig, shift_w2_sig, shift_w4_sig, shift_w_sum;
Tensor xr, xk, xv;
Tensor mix_r_sig, mix_k_sig, mix_v_sig;
Tensor r_pre, k_pre, v, r, k_exp;
Tensor decay_tmp, decay_delta, decay_pre, decay;
Tensor time_first_exp;
Tensor *num_states, *den_states;
Tensor *write_gates, *read_gates;
Tensor wkv, wkv_r, tm_out, x_after_tm;
Tensor x_ln2, xm, cm_mix_sig;
Tensor gate_pre, up_val, gate_silu, hidden, cm_out;
Tensor ffn_gu;
} LayerCache;
typedef struct {
int seq_len, n_layers;
Tensor emb_out, x_ln0;
LayerCache *layers;
Tensor x_final, x_ln_out, logits;
} ForwardCache;
static void init_layer_grads(LayerGrads *lg, const lrnnConfig *cfg, const LayerParams *lp) {
int n_embd = cfg->n_embd, ffn_h = ffn_hidden(cfg), lora_rank = cfg->decay_lora_rank;
int n_slots = cfg->n_mem_slots;
lg->ln1_weight = tensor_alloc(lp->ln1_weight.rows, lp->ln1_weight.cols);
lg->ln1_bias = tensor_alloc(lp->ln1_bias.rows, lp->ln1_bias.cols);
lg->ln2_weight = tensor_alloc(lp->ln2_weight.rows, lp->ln2_weight.cols);
lg->ln2_bias = tensor_alloc(lp->ln2_bias.rows, lp->ln2_bias.cols);
lg->time_shift_w1 = tensor_alloc_1d(n_embd); lg->time_shift_w2 = tensor_alloc_1d(n_embd); lg->time_shift_w4 = tensor_alloc_1d(n_embd);
lg->time_mix_r = tensor_alloc_1d(n_embd); lg->time_mix_k = tensor_alloc_1d(n_embd); lg->time_mix_v = tensor_alloc_1d(n_embd);
lg->decay_lora_a = tensor_alloc(n_embd, lora_rank); lg->decay_lora_b = tensor_alloc(lora_rank, n_embd);
lg->decay_base = tensor_alloc_1d(n_embd); lg->time_first = tensor_alloc_1d(n_embd);
lg->Wr = tensor_alloc(n_embd, n_embd); lg->Wk = tensor_alloc(n_embd, n_embd);
lg->Wv = tensor_alloc(n_embd, n_embd); lg->Wo = tensor_alloc(n_embd, n_embd);
lg->channel_mix = tensor_alloc_1d(n_embd);
lg->ffn_gate_up = tensor_alloc(n_embd, 2 * ffn_h); lg->ffn_down = tensor_alloc(ffn_h, n_embd);
lg->mem_gate_write = tensor_alloc(n_embd, n_slots); lg->mem_gate_read = tensor_alloc(n_embd, n_slots);
}
static void zero_layer_grads(LayerGrads *lg) {
tensor_zero(&lg->ln1_weight); tensor_zero(&lg->ln1_bias); tensor_zero(&lg->ln2_weight); tensor_zero(&lg->ln2_bias);
tensor_zero(&lg->time_shift_w1); tensor_zero(&lg->time_shift_w2); tensor_zero(&lg->time_shift_w4);
tensor_zero(&lg->time_mix_r); tensor_zero(&lg->time_mix_k); tensor_zero(&lg->time_mix_v);
tensor_zero(&lg->decay_lora_a); tensor_zero(&lg->decay_lora_b); tensor_zero(&lg->decay_base); tensor_zero(&lg->time_first);
tensor_zero(&lg->Wr); tensor_zero(&lg->Wk); tensor_zero(&lg->Wv); tensor_zero(&lg->Wo);
tensor_zero(&lg->channel_mix); tensor_zero(&lg->ffn_gate_up); tensor_zero(&lg->ffn_down);
tensor_zero(&lg->mem_gate_write); tensor_zero(&lg->mem_gate_read);
}
static void free_layer_grads(LayerGrads *lg) {
tensor_free(&lg->ln1_weight); tensor_free(&lg->ln1_bias); tensor_free(&lg->ln2_weight); tensor_free(&lg->ln2_bias);
tensor_free(&lg->time_shift_w1); tensor_free(&lg->time_shift_w2); tensor_free(&lg->time_shift_w4);
tensor_free(&lg->time_mix_r); tensor_free(&lg->time_mix_k); tensor_free(&lg->time_mix_v);
tensor_free(&lg->decay_lora_a); tensor_free(&lg->decay_lora_b); tensor_free(&lg->decay_base); tensor_free(&lg->time_first);
tensor_free(&lg->Wr); tensor_free(&lg->Wk); tensor_free(&lg->Wv); tensor_free(&lg->Wo);
tensor_free(&lg->channel_mix); tensor_free(&lg->ffn_gate_up); tensor_free(&lg->ffn_down);
tensor_free(&lg->mem_gate_write); tensor_free(&lg->mem_gate_read);
}
static void init_model_grads(ModelGrads *mg, const ModelParams *mp, const lrnnConfig *cfg) {
mg->n_layers = cfg->n_layer;
mg->emb = tensor_alloc(cfg->vocab_size, cfg->n_embd);
mg->ln0_weight = tensor_alloc_1d(cfg->n_embd); mg->ln0_bias = tensor_alloc_1d(cfg->n_embd);
mg->layers = (LayerGrads *)calloc((size_t)cfg->n_layer, sizeof(LayerGrads));
for (int i = 0; i < cfg->n_layer; i++) init_layer_grads(&mg->layers[i], cfg, &mp->layers[i]);
mg->ln_out_weight = tensor_alloc_1d(cfg->n_embd); mg->ln_out_bias = tensor_alloc_1d(cfg->n_embd);
mg->head = tensor_alloc(cfg->n_embd, cfg->vocab_size);
}
static void zero_model_grads(ModelGrads *mg) {
tensor_zero(&mg->emb); tensor_zero(&mg->ln0_weight); tensor_zero(&mg->ln0_bias);
for (int i = 0; i < mg->n_layers; i++) zero_layer_grads(&mg->layers[i]);
tensor_zero(&mg->ln_out_weight); tensor_zero(&mg->ln_out_bias); tensor_zero(&mg->head);
}
static void free_model_grads(ModelGrads *mg) {
tensor_free(&mg->emb); tensor_free(&mg->ln0_weight); tensor_free(&mg->ln0_bias);
for (int i = 0; i < mg->n_layers; i++) free_layer_grads(&mg->layers[i]);
free(mg->layers); mg->layers = NULL;
tensor_free(&mg->ln_out_weight); tensor_free(&mg->ln_out_bias); tensor_free(&mg->head);
}
static void init_layer_cache(LayerCache *lc, int seq_len, const lrnnConfig *cfg) {
if (seq_len <= 0) { fprintf(stderr, "FATAL: init_layer_cache seq_len=%d\n", seq_len); exit(1); }
int n_embd = cfg->n_embd, ffn_h = ffn_hidden(cfg), lora_rank = cfg->decay_lora_rank, n_slots = cfg->n_mem_slots;
lc->x_in = tensor_alloc(seq_len, n_embd);
lc->x_ln1 = tensor_alloc(seq_len, n_embd); lc->x_shifted = tensor_alloc(seq_len, n_embd);
lc->shift_w1_sig = tensor_alloc_1d(n_embd); lc->shift_w2_sig = tensor_alloc_1d(n_embd);
lc->shift_w4_sig = tensor_alloc_1d(n_embd); lc->shift_w_sum = tensor_alloc_1d(n_embd);
lc->xr = tensor_alloc(seq_len, n_embd); lc->xk = tensor_alloc(seq_len, n_embd); lc->xv = tensor_alloc(seq_len, n_embd);
lc->mix_r_sig = tensor_alloc_1d(n_embd); lc->mix_k_sig = tensor_alloc_1d(n_embd); lc->mix_v_sig = tensor_alloc_1d(n_embd);
lc->r_pre = tensor_alloc(seq_len, n_embd); lc->k_pre = tensor_alloc(seq_len, n_embd);
lc->v = tensor_alloc(seq_len, n_embd); lc->r = tensor_alloc(seq_len, n_embd); lc->k_exp = tensor_alloc(seq_len, n_embd);
lc->decay_tmp = tensor_alloc(seq_len, lora_rank); lc->decay_delta = tensor_alloc(seq_len, n_embd);
lc->decay_pre = tensor_alloc(seq_len, n_embd); lc->decay = tensor_alloc(seq_len, n_embd);
lc->time_first_exp = tensor_alloc_1d(n_embd);
lc->wkv = tensor_alloc(seq_len, n_embd); lc->wkv_r = tensor_alloc(seq_len, n_embd);
lc->tm_out = tensor_alloc(seq_len, n_embd); lc->x_after_tm = tensor_alloc(seq_len, n_embd);
lc->x_ln2 = tensor_alloc(seq_len, n_embd); lc->xm = tensor_alloc(seq_len, n_embd);
lc->cm_mix_sig = tensor_alloc_1d(n_embd);
lc->gate_pre = tensor_alloc(seq_len, ffn_h); lc->up_val = tensor_alloc(seq_len, ffn_h);
lc->ffn_gu = tensor_alloc(seq_len, 2 * ffn_h);
lc->gate_silu = tensor_alloc(seq_len, ffn_h); lc->hidden = tensor_alloc(seq_len, ffn_h);
lc->cm_out = tensor_alloc(seq_len, n_embd);
lc->num_states = (Tensor *)calloc((size_t)(seq_len + 1), sizeof(Tensor));
lc->den_states = (Tensor *)calloc((size_t)(seq_len + 1), sizeof(Tensor));
for (int t = 0; t <= seq_len; t++) { lc->num_states[t] = tensor_alloc(n_slots, n_embd); lc->den_states[t] = tensor_alloc(n_slots, n_embd); }
lc->write_gates = (Tensor *)calloc((size_t)seq_len, sizeof(Tensor));
lc->read_gates = (Tensor *)calloc((size_t)seq_len, sizeof(Tensor));
for (int t = 0; t < seq_len; t++) { lc->write_gates[t] = tensor_alloc_1d(n_slots); lc->read_gates[t] = tensor_alloc_1d(n_slots); }
}
static void free_layer_cache(LayerCache *lc, int seq_len) {
tensor_free(&lc->x_in); tensor_free(&lc->x_ln1); tensor_free(&lc->x_shifted);
tensor_free(&lc->shift_w1_sig); tensor_free(&lc->shift_w2_sig); tensor_free(&lc->shift_w4_sig); tensor_free(&lc->shift_w_sum);
tensor_free(&lc->xr); tensor_free(&lc->xk); tensor_free(&lc->xv);
tensor_free(&lc->mix_r_sig); tensor_free(&lc->mix_k_sig); tensor_free(&lc->mix_v_sig);
tensor_free(&lc->r_pre); tensor_free(&lc->k_pre); tensor_free(&lc->v); tensor_free(&lc->r); tensor_free(&lc->k_exp);
tensor_free(&lc->decay_tmp); tensor_free(&lc->decay_delta); tensor_free(&lc->decay_pre); tensor_free(&lc->decay);
tensor_free(&lc->time_first_exp);
for (int t = 0; t <= seq_len; t++) { tensor_free(&lc->num_states[t]); tensor_free(&lc->den_states[t]); }
free(lc->num_states); free(lc->den_states);
tensor_free(&lc->wkv); tensor_free(&lc->wkv_r); tensor_free(&lc->tm_out); tensor_free(&lc->x_after_tm);
tensor_free(&lc->x_ln2); tensor_free(&lc->xm); tensor_free(&lc->cm_mix_sig);
tensor_free(&lc->gate_pre); tensor_free(&lc->up_val); tensor_free(&lc->ffn_gu); tensor_free(&lc->gate_silu); tensor_free(&lc->hidden); tensor_free(&lc->cm_out);
for (int t = 0; t < seq_len; t++) { tensor_free(&lc->write_gates[t]); tensor_free(&lc->read_gates[t]); }
free(lc->write_gates); free(lc->read_gates);
}
static void init_forward_cache(ForwardCache *fc, int seq_len, const lrnnConfig *cfg) {
fc->seq_len = seq_len; fc->n_layers = cfg->n_layer;
fc->emb_out = tensor_alloc(seq_len, cfg->n_embd); fc->x_ln0 = tensor_alloc(seq_len, cfg->n_embd);
fc->layers = (LayerCache *)calloc((size_t)cfg->n_layer, sizeof(LayerCache));
for (int i = 0; i < cfg->n_layer; i++) init_layer_cache(&fc->layers[i], seq_len, cfg);
fc->x_final = tensor_alloc(seq_len, cfg->n_embd); fc->x_ln_out = tensor_alloc(seq_len, cfg->n_embd);
fc->logits = tensor_alloc(seq_len, cfg->vocab_size);
}
static void free_forward_cache(ForwardCache *fc) {
tensor_free(&fc->emb_out); tensor_free(&fc->x_ln0);
for (int i = 0; i < fc->n_layers; i++) free_layer_cache(&fc->layers[i], fc->seq_len);
free(fc->layers);
tensor_free(&fc->x_final); tensor_free(&fc->x_ln_out); tensor_free(&fc->logits);
}
static void sigmoid_backward(float *d_input, const float *d_output, const float *y, int n) {
for (int i = 0; i < n; i++) d_input[i] = d_output[i] * y[i] * (1.0f - y[i]);
}
static void silu_backward(float *d_input, const float *d_output, const float *x, int n) {
for (int i = 0; i < n; i++) { float s = sigmoid_f(x[i]); d_input[i] = d_output[i] * s * (1.0f + x[i] * (1.0f - s)); }
}
static void exp_backward_clamped(float *d_input, const float *d_output, const float *x, int n) {
int i = 0;
for (; i + 8 <= n; i += 8) {
__m256 xv = _mm256_loadu_ps(x + i);
__m256 inr = _mm256_and_ps(_mm256_cmp_ps(xv, _mm256_set1_ps(-10.0f), _CMP_GE_OQ),
_mm256_cmp_ps(xv, _mm256_set1_ps(10.0f), _CMP_LE_OQ));
__m256 cv = _mm256_min_ps(_mm256_max_ps(xv, _mm256_set1_ps(-10.0f)), _mm256_set1_ps(10.0f));
__m256 e = _mm256_and_ps(exp256_ps(cv), inr);
_mm256_storeu_ps(d_input + i, _mm256_mul_ps(_mm256_loadu_ps(d_output + i), e));
}
for (; i < n; i++) {
if (x[i] < -10.0f || x[i] > 10.0f) d_input[i] = 0.0f;
else d_input[i] = d_output[i] * expf(clamp_f(x[i], -10.0f, 10.0f));
}
}
static float *backward_gemm_scratch(size_t need) {
static float *s = NULL;
static size_t cap = 0;
if (need > cap) { free(s); s = (float *)malloc(need); cap = need; }
return s;
}
static void transpose_mat(float *dst, const float *src, int rows, int cols) {
int i = 0;
for (; i + 8 <= rows; i += 8) {
int j = 0;
for (; j + 8 <= cols; j += 8) {
__m256 r0 = _mm256_loadu_ps(src + (long)(i + 0) * cols + j);
__m256 r1 = _mm256_loadu_ps(src + (long)(i + 1) * cols + j);
__m256 r2 = _mm256_loadu_ps(src + (long)(i + 2) * cols + j);
__m256 r3 = _mm256_loadu_ps(src + (long)(i + 3) * cols + j);
__m256 r4 = _mm256_loadu_ps(src + (long)(i + 4) * cols + j);
__m256 r5 = _mm256_loadu_ps(src + (long)(i + 5) * cols + j);
__m256 r6 = _mm256_loadu_ps(src + (long)(i + 6) * cols + j);
__m256 r7 = _mm256_loadu_ps(src + (long)(i + 7) * cols + j);
__m256 t0 = _mm256_unpacklo_ps(r0, r2), t1 = _mm256_unpackhi_ps(r0, r2);
__m256 t2 = _mm256_unpacklo_ps(r1, r3), t3 = _mm256_unpackhi_ps(r1, r3);
__m256 t4 = _mm256_unpacklo_ps(r4, r6), t5 = _mm256_unpackhi_ps(r4, r6);
__m256 t6 = _mm256_unpacklo_ps(r5, r7), t7 = _mm256_unpackhi_ps(r5, r7);
__m256 u0 = _mm256_unpacklo_ps(t0, t2), u1 = _mm256_unpackhi_ps(t0, t2);
__m256 u2 = _mm256_unpacklo_ps(t1, t3), u3 = _mm256_unpackhi_ps(t1, t3);
__m256 u4 = _mm256_unpacklo_ps(t4, t6), u5 = _mm256_unpackhi_ps(t4, t6);
__m256 u6 = _mm256_unpacklo_ps(t5, t7), u7 = _mm256_unpackhi_ps(t5, t7);
__m256 v0 = _mm256_permute2f128_ps(u0, u4, 0x20), v1 = _mm256_permute2f128_ps(u1, u5, 0x20);
__m256 v2 = _mm256_permute2f128_ps(u2, u6, 0x20), v3 = _mm256_permute2f128_ps(u3, u7, 0x20);
__m256 v4 = _mm256_permute2f128_ps(u0, u4, 0x31), v5 = _mm256_permute2f128_ps(u1, u5, 0x31);
__m256 v6 = _mm256_permute2f128_ps(u2, u6, 0x31), v7 = _mm256_permute2f128_ps(u3, u7, 0x31);
_mm256_storeu_ps(dst + (long)(j + 0) * rows + i, v0);
_mm256_storeu_ps(dst + (long)(j + 1) * rows + i, v1);
_mm256_storeu_ps(dst + (long)(j + 2) * rows + i, v2);
_mm256_storeu_ps(dst + (long)(j + 3) * rows + i, v3);
_mm256_storeu_ps(dst + (long)(j + 4) * rows + i, v4);
_mm256_storeu_ps(dst + (long)(j + 5) * rows + i, v5);
_mm256_storeu_ps(dst + (long)(j + 6) * rows + i, v6);
_mm256_storeu_ps(dst + (long)(j + 7) * rows + i, v7);
}
for (; j < cols; j++)
for (int ii = 0; ii < 8; ii++)
dst[(long)j * rows + i + ii] = src[(long)(i + ii) * cols + j];
}
for (; i < rows; i++)
for (int j = 0; j < cols; j++)
dst[(long)j * rows + i] = src[(long)i * cols + j];
}
static struct {
float *buf;
size_t cap, top;
} g_bwd_arena;
static void bwd_reset(void) { g_bwd_arena.top = 0; }
static void *bwd_alloc(size_t n) {
n = (n + 31u) & ~(size_t)31u;
if (g_bwd_arena.top + n > g_bwd_arena.cap) {
size_t nc = g_bwd_arena.cap ? g_bwd_arena.cap * 2 : (size_t)1 << 22;
while (nc < g_bwd_arena.top + n) nc *= 2;
g_bwd_arena.buf = (float *)realloc(g_bwd_arena.buf, nc);
if (!g_bwd_arena.buf) { fprintf(stderr, "FATAL: bwd arena realloc %zu\n", nc); exit(1); }
g_bwd_arena.cap = nc;
}
void *p = (char *)g_bwd_arena.buf + g_bwd_arena.top;
g_bwd_arena.top += n;
return p;
}
static Tensor bwd_tensor(int rows, int cols) {
Tensor t = { .rows = rows, .cols = cols, .size = rows * cols, .data = NULL };
if (t.size > 0) t.data = (float *)bwd_alloc((size_t)t.size * sizeof(float));
return t;
}
#define BWD_FREE(t) ((void)0)
static void matmul_backward_x(Tensor *d_X, const Tensor *d_Y, const Tensor *W) {
int seq_len = d_Y->rows, out_dim = d_Y->cols, in_dim = W->rows;
long flops = (long)seq_len * in_dim * out_dim;
if (flops < 100000) {
for (int s = 0; s < seq_len; s++) {
for (int i = 0; i < in_dim; i++) {
float sum = 0.0f;
for (int j = 0; j < out_dim; j++) sum += d_Y->data[s * out_dim + j] * W->data[i * out_dim + j];
d_X->data[s * in_dim + i] = sum;
}
}
return;
}
float *base = backward_gemm_scratch((size_t)in_dim * out_dim * sizeof(float));
transpose_mat(base, W->data, in_dim, out_dim);
hp_sgemm(d_X->data, d_Y->data, base, seq_len, in_dim, out_dim);
}
static void matmul_backward_w(Tensor *d_W, const Tensor *d_Y, const Tensor *X) {
int seq_len = X->rows, in_dim = X->cols, out_dim = d_Y->cols;
long flops = (long)seq_len * in_dim * out_dim;
if (flops < 100000) {
for (int i = 0; i < in_dim; i++) {
for (int j = 0; j < out_dim; j++) {
float sum = 0.0f;
for (int s = 0; s < seq_len; s++) sum += X->data[s * in_dim + i] * d_Y->data[s * out_dim + j];
d_W->data[i * out_dim + j] += sum;
}
}
return;
}
size_t a = (size_t)in_dim * seq_len;
size_t b = (size_t)in_dim * out_dim;
float *base = backward_gemm_scratch((a + b) * sizeof(float));
float *Xt = base, *tmp = base + a;
transpose_mat(Xt, X->data, seq_len, in_dim);
hp_sgemm(tmp, Xt, d_Y->data, in_dim, out_dim, seq_len);
for (int i = 0; i < d_W->size; i++) d_W->data[i] += tmp[i];
}
static void layer_norm_backward_single(float *d_x, float *d_weight, float *d_bias,
const float *d_y, const float *x,
const float *weight, int n) {
__m256 acc_mean = _mm256_setzero_ps();
int i = 0;
for (; i + 8 <= n; i += 8) acc_mean = _mm256_add_ps(acc_mean, _mm256_loadu_ps(x + i));
__m128 lo = _mm_add_ps(_mm256_castps256_ps128(acc_mean), _mm256_extractf128_ps(acc_mean, 1));
lo = _mm_hadd_ps(lo, lo); lo = _mm_hadd_ps(lo, lo);
float mean = _mm_cvtss_f32(lo);
for (; i < n; i++) mean += x[i];
mean /= (float)n;
__m256 vmean = _mm256_set1_ps(mean);
__m256 acc_var = _mm256_setzero_ps();
i = 0;
for (; i + 8 <= n; i += 8) {
__m256 d = _mm256_sub_ps(_mm256_loadu_ps(x + i), vmean);
acc_var = _mm256_fmadd_ps(d, d, acc_var);
}
lo = _mm_add_ps(_mm256_castps256_ps128(acc_var), _mm256_extractf128_ps(acc_var, 1));
lo = _mm_hadd_ps(lo, lo); lo = _mm_hadd_ps(lo, lo);
float var = _mm_cvtss_f32(lo);
for (; i < n; i++) { float d = x[i] - mean; var += d * d; }
var /= (float)n;
float inv_std = 1.0f / sqrtf(var + EPSILON);
__m256 vinv_std = _mm256_set1_ps(inv_std);
__m256 acc1 = _mm256_setzero_ps(), acc2 = _mm256_setzero_ps();
i = 0;
for (; i + 8 <= n; i += 8) {
__m256 vx = _mm256_loadu_ps(x + i);
__m256 xh = _mm256_mul_ps(_mm256_sub_ps(vx, vmean), vinv_std);
__m256 dy = _mm256_loadu_ps(d_y + i);
__m256 vw = _mm256_loadu_ps(weight + i);
__m256 dg = _mm256_mul_ps(dy, vw);
acc1 = _mm256_add_ps(acc1, dg);
acc2 = _mm256_add_ps(acc2, _mm256_mul_ps(dg, xh));
}
__m256 v1 = _mm256_add_ps(acc1, _mm256_permute2f128_ps(acc1, acc1, 1));
v1 = _mm256_hadd_ps(v1, v1); v1 = _mm256_hadd_ps(v1, v1);
float sum_dy_gamma = _mm_cvtss_f32(_mm256_castps256_ps128(v1));
__m256 v2 = _mm256_add_ps(acc2, _mm256_permute2f128_ps(acc2, acc2, 1));
v2 = _mm256_hadd_ps(v2, v2); v2 = _mm256_hadd_ps(v2, v2);
float sum_dy_gamma_xhat = _mm_cvtss_f32(_mm256_castps256_ps128(v2));
for (; i < n; i++) {
float x_hat = (x[i] - mean) * inv_std;
float dy_gamma = d_y[i] * weight[i];
sum_dy_gamma += dy_gamma;
sum_dy_gamma_xhat += dy_gamma * x_hat;
}
float inv_n = 1.0f / (float)n;
__m256 vis = _mm256_set1_ps(inv_n * inv_std);
__m256 vs1 = _mm256_set1_ps(sum_dy_gamma), vs2 = _mm256_set1_ps(sum_dy_gamma_xhat);
__m256 vn = _mm256_set1_ps((float)n);
i = 0;
for (; i + 8 <= n; i += 8) {
__m256 vx = _mm256_loadu_ps(x + i);
__m256 xh = _mm256_mul_ps(_mm256_sub_ps(vx, vmean), vinv_std);
__m256 dy = _mm256_loadu_ps(d_y + i);
__m256 vw = _mm256_loadu_ps(weight + i);
__m256 dg = _mm256_mul_ps(dy, vw);
__m256 dx = _mm256_mul_ps(vis, _mm256_sub_ps(_mm256_sub_ps(_mm256_mul_ps(vn, dg), vs1), _mm256_mul_ps(xh, vs2)));
_mm256_storeu_ps(d_x + i, dx);
_mm256_storeu_ps(d_weight + i, _mm256_add_ps(_mm256_loadu_ps(d_weight + i), _mm256_mul_ps(dy, xh)));
_mm256_storeu_ps(d_bias + i, _mm256_add_ps(_mm256_loadu_ps(d_bias + i), dy));
}
for (; i < n; i++) {
float x_hat = (x[i] - mean) * inv_std;
d_weight[i] += d_y[i] * x_hat;
d_bias[i] += d_y[i];
float dy_gamma = d_y[i] * weight[i];
d_x[i] = inv_n * inv_std * ((float)n * dy_gamma - sum_dy_gamma - x_hat * sum_dy_gamma_xhat);
}
}
static void layer_norm_backward_seq(Tensor *d_x, Tensor *d_weight, Tensor *d_bias,
const Tensor *d_y, const Tensor *x, const Tensor *weight) {
int seq_len = x->rows, n_embd = x->cols;
for (int s = 0; s < seq_len; s++)
layer_norm_backward_single(d_x->data + s * n_embd, d_weight->data, d_bias->data,
d_y->data + s * n_embd, x->data + s * n_embd, weight->data, n_embd);
}
static void softmax_cross_entropy_backward(Tensor *d_logits, const Tensor *logits, const int *targets) {
int seq_len = logits->rows, vocab_size = logits->cols;
float *probs = (float *)malloc((size_t)vocab_size * sizeof(float));
for (int t = 0; t < seq_len; t++) {
softmax_vec(probs, logits->data + t * vocab_size, vocab_size);
int target = targets[t];
for (int v = 0; v < vocab_size; v++) {
float q = (v == target) ? 1.0f : 0.0f;
d_logits->data[t * vocab_size + v] = (probs[v] - q) / (float)seq_len;
}
}
free(probs);
}
static void wkv_backward(
Tensor *d_k_exp, Tensor *d_v, Tensor *d_decay, Tensor *d_time_first_exp,
Tensor *d_mem_gate_write_proj, Tensor *d_mem_gate_read_proj,
const Tensor *d_wkv, const Tensor *k_exp, const Tensor *v,
const Tensor *decay, const Tensor *time_first_exp,
Tensor *num_states, Tensor *den_states,
Tensor *write_gates, Tensor *read_gates,
int seq_len, int n_embd, int n_head, int n_slots, const float *alibi_slopes) {
int hdim = n_embd / n_head;
tensor_zero(d_k_exp); tensor_zero(d_v); tensor_zero(d_decay); tensor_zero(d_time_first_exp);
float *d_num_next = (float *)bwd_alloc((size_t)(n_slots * n_embd) * sizeof(float));
float *d_den_next = (float *)bwd_alloc((size_t)(n_slots * n_embd) * sizeof(float));
memset(d_num_next, 0, (size_t)(n_slots * n_embd) * sizeof(float));
memset(d_den_next, 0, (size_t)(n_slots * n_embd) * sizeof(float));
float *d_write_g = (float *)bwd_alloc((size_t)n_slots * sizeof(float));
float *d_read_g = (float *)bwd_alloc((size_t)n_slots * sizeof(float));
float *alibi_decay_arr = (float *)bwd_alloc((size_t)n_head * sizeof(float));
for (int h = 0; h < n_head; h++) alibi_decay_arr[h] = expf(-alibi_slopes[h]);
for (int t = seq_len - 1; t >= 0; t--) {
memset(d_write_g, 0, (size_t)n_slots * sizeof(float));
memset(d_read_g, 0, (size_t)n_slots * sizeof(float));
for (int h = 0; h < n_head; h++) {
int base_h = h * hdim; float ad = alibi_decay_arr[h];
for (int d_idx = 0; d_idx < hdim; d_idx++) {
int i = base_h + d_idx, idx = t * n_embd + i;
float ki = k_exp->data[idx], vi = v->data[idx], di = decay->data[idx];
float tfi = time_first_exp->data[i], kv = ki * vi, combined = di * ad;
float read_num = 0.0f, read_den = 0.0f;
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
read_num += read_gates[t].data[s] * num_states[t].data[si];
read_den += read_gates[t].data[s] * den_states[t].data[si];
}
float numerator = read_num + tfi * kv, denominator = read_den + tfi * ki + EPSILON;
float inv_den = 1.0f / denominator;
float dw = d_wkv->data[idx];
float d_numerator = dw * inv_den, d_denominator = -dw * numerator * inv_den * inv_den;
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
d_read_g[s] += d_numerator * num_states[t].data[si] + d_denominator * den_states[t].data[si];
float d_state_num = d_numerator * read_gates[t].data[s] + d_num_next[si] * combined;
float d_state_den = d_denominator * read_gates[t].data[s] + d_den_next[si] * combined;
d_decay->data[idx] += (d_num_next[si] * num_states[t].data[si] + d_den_next[si] * den_states[t].data[si]) * ad;
d_write_g[s] += d_num_next[si] * kv + d_den_next[si] * ki;
float wg = write_gates[t].data[s];
d_k_exp->data[idx] += d_num_next[si] * wg * vi + d_den_next[si] * wg;
d_v->data[idx] += d_num_next[si] * wg * ki;
d_num_next[si] = d_state_num; d_den_next[si] = d_state_den;
}
d_time_first_exp->data[i] += d_numerator * kv + d_denominator * ki;
d_k_exp->data[idx] += d_numerator * tfi * vi + d_denominator * tfi;
d_v->data[idx] += d_numerator * tfi * ki;
}
}
{
float *wg = write_gates[t].data, *rg = read_gates[t].data;
float dot_w = 0.0f, dot_r = 0.0f;
for (int s = 0; s < n_slots; s++) { dot_w += d_write_g[s] * wg[s]; dot_r += d_read_g[s] * rg[s]; }
for (int s = 0; s < n_slots; s++) {
d_mem_gate_write_proj->data[t * n_slots + s] = wg[s] * (d_write_g[s] - dot_w);
d_mem_gate_read_proj->data[t * n_slots + s] = rg[s] * (d_read_g[s] - dot_r);
}
}
}
for (int i = 0; i < d_k_exp->size; i++) { d_k_exp->data[i] = clamp_f(d_k_exp->data[i], -GRAD_CLIP, GRAD_CLIP); d_v->data[i] = clamp_f(d_v->data[i], -GRAD_CLIP, GRAD_CLIP); }
for (int i = 0; i < d_decay->size; i++) d_decay->data[i] = clamp_f(d_decay->data[i], -GRAD_CLIP, GRAD_CLIP);
for (int i = 0; i < n_embd; i++) d_time_first_exp->data[i] = clamp_f(d_time_first_exp->data[i], -GRAD_CLIP, GRAD_CLIP);
}
static void multi_scale_shift_backward(Tensor *d_x, Tensor *d_shift_w1, Tensor *d_shift_w2, Tensor *d_shift_w4,
const Tensor *d_out, const Tensor *x, const Tensor *shift_w1_sig, const Tensor *shift_w2_sig,
const Tensor *shift_w4_sig, const Tensor *shift_w_sum, int seq_len, int n_embd) {
tensor_zero(d_x); tensor_zero(d_shift_w1); tensor_zero(d_shift_w2); tensor_zero(d_shift_w4);
for (int t = 0; t < seq_len; t++) for (int i = 0; i < n_embd; i++) {
float w1 = shift_w1_sig->data[i], w2 = shift_w2_sig->data[i], w4 = shift_w4_sig->data[i];
float inv_sum = 1.0f / shift_w_sum->data[i];
float x1 = (t >= 1) ? x->data[(t-1) * n_embd + i] : 0.0f;
float x2 = (t >= 2) ? x->data[(t-2) * n_embd + i] : 0.0f;
float x4 = (t >= 4) ? x->data[(t-4) * n_embd + i] : 0.0f;
float d_out_ti = d_out->data[t * n_embd + i];
if (t >= 1) d_x->data[(t-1) * n_embd + i] += d_out_ti * w1 * inv_sum;
if (t >= 2) d_x->data[(t-2) * n_embd + i] += d_out_ti * w2 * inv_sum;
if (t >= 4) d_x->data[(t-4) * n_embd + i] += d_out_ti * w4 * inv_sum;
float numerator = w1 * x1 + w2 * x2 + w4 * x4;
d_shift_w1->data[i] += d_out_ti * (x1 * inv_sum - numerator * inv_sum * inv_sum) * w1 * (1.0f - w1);
d_shift_w2->data[i] += d_out_ti * (x2 * inv_sum - numerator * inv_sum * inv_sum) * w2 * (1.0f - w2);
d_shift_w4->data[i] += d_out_ti * (x4 * inv_sum - numerator * inv_sum * inv_sum) * w4 * (1.0f - w4);
}
}
static void token_mixing_backward(Tensor *d_x_ln1, Tensor *d_x_shifted, Tensor *d_mix_r, Tensor *d_mix_k, Tensor *d_mix_v,
const Tensor *d_xr, const Tensor *d_xk, const Tensor *d_xv,
const Tensor *x_ln1, const Tensor *x_shifted, const Tensor *mix_r_sig, const Tensor *mix_k_sig, const Tensor *mix_v_sig,
int seq_len, int n_embd) {
tensor_zero(d_x_ln1); tensor_zero(d_x_shifted); tensor_zero(d_mix_r); tensor_zero(d_mix_k); tensor_zero(d_mix_v);
for (int t = 0; t < seq_len; t++) for (int i = 0; i < n_embd; i++) {
int idx = t * n_embd + i;
float mr = mix_r_sig->data[i], mk = mix_k_sig->data[i], mv = mix_v_sig->data[i];
float x_val = x_ln1->data[idx], x_sh = x_shifted->data[idx];
d_x_ln1->data[idx] += d_xr->data[idx] * mr + d_xk->data[idx] * mk + d_xv->data[idx] * mv;
d_x_shifted->data[idx] += d_xr->data[idx] * (1.0f - mr) + d_xk->data[idx] * (1.0f - mk) + d_xv->data[idx] * (1.0f - mv);
d_mix_r->data[i] += d_xr->data[idx] * (x_val - x_sh) * mr * (1.0f - mr);
d_mix_k->data[i] += d_xk->data[idx] * (x_val - x_sh) * mk * (1.0f - mk);
d_mix_v->data[i] += d_xv->data[idx] * (x_val - x_sh) * mv * (1.0f - mv);
}
}
static void channel_mix_shift_backward(Tensor *d_x_ln2, Tensor *d_channel_mix, const Tensor *d_xm,
const Tensor *x_ln2, const Tensor *cm_mix_sig, int seq_len, int n_embd) {
tensor_zero(d_x_ln2); tensor_zero(d_channel_mix);
for (int t = 0; t < seq_len; t++) for (int i = 0; i < n_embd; i++) {
int idx = t * n_embd + i;
float mix = cm_mix_sig->data[i];
float x_curr = x_ln2->data[idx], x_prev = (t > 0) ? x_ln2->data[(t-1) * n_embd + i] : 0.0f;
d_x_ln2->data[idx] += d_xm->data[idx] * mix;
if (t > 0) d_x_ln2->data[(t-1) * n_embd + i] += d_xm->data[idx] * (1.0f - mix);
d_channel_mix->data[i] += d_xm->data[idx] * (x_curr - x_prev) * mix * (1.0f - mix);
}
}
static float forward_with_cache(ForwardCache *cache, const int *tokens, int seq_len,
const ModelParams *mp, const lrnnConfig *cfg) {
int n_embd = cfg->n_embd, n_head = cfg->n_head, ffn_h = ffn_hidden(cfg);
for (int t = 0; t < seq_len; t++)
memcpy(cache->emb_out.data + t * n_embd, mp->emb.data + tokens[t] * n_embd, (size_t)n_embd * sizeof(float));
layer_norm_seq(&cache->x_ln0, &cache->emb_out, &mp->ln0_weight, &mp->ln0_bias);
tensor_copy(&cache->x_final, &cache->x_ln0);
for (int layer_idx = 0; layer_idx < mp->n_layers; layer_idx++) {
const LayerParams *lp = &mp->layers[layer_idx];
LayerCache *lc = &cache->layers[layer_idx];
tensor_copy(&lc->x_in, &cache->x_final);
layer_norm_seq(&lc->x_ln1, &cache->x_final, &lp->ln1_weight, &lp->ln1_bias);
sigmoid_vec(lc->shift_w1_sig.data, lp->time_shift_w1.data, n_embd);
sigmoid_vec(lc->shift_w2_sig.data, lp->time_shift_w2.data, n_embd);
sigmoid_vec(lc->shift_w4_sig.data, lp->time_shift_w4.data, n_embd);
for (int i = 0; i < n_embd; i++) lc->shift_w_sum.data[i] = lc->shift_w1_sig.data[i] + lc->shift_w2_sig.data[i] + lc->shift_w4_sig.data[i] + EPSILON;
sigmoid_vec(lc->mix_r_sig.data, lp->time_mix_r.data, n_embd);
sigmoid_vec(lc->mix_k_sig.data, lp->time_mix_k.data, n_embd);
sigmoid_vec(lc->mix_v_sig.data, lp->time_mix_v.data, n_embd);
for (int t = 0; t < seq_len; t++) {
const float *x1p = (t >= 1) ? lc->x_ln1.data + (long)(t - 1) * n_embd : NULL;
const float *x2p = (t >= 2) ? lc->x_ln1.data + (long)(t - 2) * n_embd : NULL;
const float *x4p = (t >= 4) ? lc->x_ln1.data + (long)(t - 4) * n_embd : NULL;
__m256 vzero = _mm256_setzero_ps();
int i = 0;
for (; i + 8 <= n_embd; i += 8) {
__m256 s1 = _mm256_loadu_ps(lc->shift_w1_sig.data + i);
__m256 s2 = _mm256_loadu_ps(lc->shift_w2_sig.data + i);
__m256 s4 = _mm256_loadu_ps(lc->shift_w4_sig.data + i);
__m256 isum = _mm256_div_ps(_mm256_set1_ps(1.0f), _mm256_loadu_ps(lc->shift_w_sum.data + i));
__m256 w1 = _mm256_mul_ps(s1, isum), w2 = _mm256_mul_ps(s2, isum), w4 = _mm256_mul_ps(s4, isum);
__m256 x1 = x1p ? _mm256_loadu_ps(x1p + i) : vzero;
__m256 x2 = x2p ? _mm256_loadu_ps(x2p + i) : vzero;
__m256 x4 = x4p ? _mm256_loadu_ps(x4p + i) : vzero;
__m256 xs = _mm256_fmadd_ps(w1, x1, _mm256_fmadd_ps(w2, x2, _mm256_mul_ps(w4, x4)));
_mm256_storeu_ps(lc->x_shifted.data + (long)t * n_embd + i, xs);
__m256 xn = _mm256_loadu_ps(lc->x_ln1.data + (long)t * n_embd + i);
__m256 mr = _mm256_loadu_ps(lc->mix_r_sig.data + i);
__m256 mk = _mm256_loadu_ps(lc->mix_k_sig.data + i);
__m256 mv = _mm256_loadu_ps(lc->mix_v_sig.data + i);
__m256 one = _mm256_set1_ps(1.0f);
_mm256_storeu_ps(lc->xr.data + (long)t * n_embd + i, _mm256_fmadd_ps(xn, mr, _mm256_mul_ps(xs, _mm256_sub_ps(one, mr))));
_mm256_storeu_ps(lc->xk.data + (long)t * n_embd + i, _mm256_fmadd_ps(xn, mk, _mm256_mul_ps(xs, _mm256_sub_ps(one, mk))));
_mm256_storeu_ps(lc->xv.data + (long)t * n_embd + i, _mm256_fmadd_ps(xn, mv, _mm256_mul_ps(xs, _mm256_sub_ps(one, mv))));
}
for (; i < n_embd; i++) {
float w1 = lc->shift_w1_sig.data[i] / lc->shift_w_sum.data[i];
float w2 = lc->shift_w2_sig.data[i] / lc->shift_w_sum.data[i];
float w4 = lc->shift_w4_sig.data[i] / lc->shift_w_sum.data[i];
float x1 = (t >= 1) ? lc->x_ln1.data[(t-1) * n_embd + i] : 0.0f;
float x2 = (t >= 2) ? lc->x_ln1.data[(t-2) * n_embd + i] : 0.0f;
float x4 = (t >= 4) ? lc->x_ln1.data[(t-4) * n_embd + i] : 0.0f;
float xs = w1 * x1 + w2 * x2 + w4 * x4;
lc->x_shifted.data[t * n_embd + i] = xs;
int idx = t * n_embd + i;
lc->xr.data[idx] = lc->x_ln1.data[idx] * lc->mix_r_sig.data[i] + xs * (1.0f - lc->mix_r_sig.data[i]);
lc->xk.data[idx] = lc->x_ln1.data[idx] * lc->mix_k_sig.data[i] + xs * (1.0f - lc->mix_k_sig.data[i]);
lc->xv.data[idx] = lc->x_ln1.data[idx] * lc->mix_v_sig.data[i] + xs * (1.0f - lc->mix_v_sig.data[i]);
}
}
matmul(&lc->r_pre, &lc->xr, &lp->Wr); matmul(&lc->k_pre, &lc->xk, &lp->Wk); matmul(&lc->v, &lc->xv, &lp->Wv);
matmul(&lc->decay_tmp, &lc->x_ln1, &lp->decay_lora_a); matmul(&lc->decay_delta, &lc->decay_tmp, &lp->decay_lora_b);
for (int t = 0; t < seq_len; t++) {
int i = 0;
for (; i + 8 <= n_embd; i += 8) {
int idx = t * n_embd + i;
__m256 dp = _mm256_add_ps(_mm256_loadu_ps(lp->decay_base.data + i), _mm256_loadu_ps(lc->decay_delta.data + idx));
_mm256_storeu_ps(lc->decay_pre.data + idx, dp);
_mm256_storeu_ps(lc->decay.data + idx, sigmoid256(dp));
}
for (; i < n_embd; i++) {
int idx = t * n_embd + i;
lc->decay_pre.data[idx] = lp->decay_base.data[i] + lc->decay_delta.data[idx];
lc->decay.data[idx] = sigmoid_f(lc->decay_pre.data[idx]);
}
}
sigmoid_vec(lc->r.data, lc->r_pre.data, lc->r_pre.size);
exp_vec(lc->time_first_exp.data, lp->time_first.data, n_embd);
exp_vec(lc->k_exp.data, lc->k_pre.data, lc->k_pre.size);
{
int n_slots = cfg->n_mem_slots, hdim_val = n_embd / n_head;
tensor_zero(&lc->num_states[0]); tensor_zero(&lc->den_states[0]);
float wl[16];
float rl[16];
for (int t = 0; t < seq_len; t++) {
matvec(wl, lc->x_ln1.data + t * n_embd, &lp->mem_gate_write);
softmax_vec(lc->write_gates[t].data, wl, n_slots);
matvec(rl, lc->x_ln1.data + t * n_embd, &lp->mem_gate_read);
softmax_vec(lc->read_gates[t].data, rl, n_slots);
for (int h = 0; h < n_head; h++) {
int base = h * hdim_val;
float ad = expf(-lp->alibi_slopes.data[h]);
__m256 vad = _mm256_set1_ps(ad), veps = _mm256_set1_ps(EPSILON);
int d = 0;
for (; d + 8 <= hdim_val; d += 8) {
int i = base + d, idx = t * n_embd + i;
__m256 ki = _mm256_loadu_ps(lc->k_exp.data + idx);
__m256 vi = _mm256_loadu_ps(lc->v.data + idx);
__m256 di = _mm256_loadu_ps(lc->decay.data + idx);
__m256 vkv = _mm256_mul_ps(ki, vi);
__m256 vco = _mm256_mul_ps(di, vad);
__m256 num = _mm256_setzero_ps(), den = _mm256_setzero_ps();
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
__m256 rgv = _mm256_set1_ps(lc->read_gates[t].data[s]);
num = _mm256_fmadd_ps(rgv, _mm256_loadu_ps(lc->num_states[t].data + si), num);
den = _mm256_fmadd_ps(rgv, _mm256_loadu_ps(lc->den_states[t].data + si), den);
}
__m256 tfi = _mm256_loadu_ps(lc->time_first_exp.data + i);
num = _mm256_fmadd_ps(tfi, vkv, num);
den = _mm256_add_ps(_mm256_fmadd_ps(tfi, ki, den), veps);
_mm256_storeu_ps(lc->wkv.data + idx, _mm256_div_ps(num, den));
for (int s = 0; s < n_slots; s++) {
int si = s * n_embd + i;
float wg = lc->write_gates[t].data[s];
__m256 vwg = _mm256_set1_ps(wg);
_mm256_storeu_ps(lc->num_states[t+1].data + si,
_mm256_fmadd_ps(vwg, vkv, _mm256_mul_ps(vco, _mm256_loadu_ps(lc->num_states[t].data + si))));
_mm256_storeu_ps(lc->den_states[t+1].data + si,
_mm256_fmadd_ps(vwg, ki, _mm256_mul_ps(vco, _mm256_loadu_ps(lc->den_states[t].data + si))));
}
}
for (; d < hdim_val; d++) {
int i = base + d, idx = t * n_embd + i;
float ki = lc->k_exp.data[idx], vi = lc->v.data[idx], di = lc->decay.data[idx];
float tfi = lc->time_first_exp.data[i], kv = ki * vi, combined = di * ad;
float rn = 0.0f, rd = 0.0f;
for (int s = 0; s < n_slots; s++) { int si = s * n_embd + i;
rn += lc->read_gates[t].data[s] * lc->num_states[t].data[si];
rd += lc->read_gates[t].data[s] * lc->den_states[t].data[si]; }
lc->wkv.data[idx] = (rn + tfi * kv) / (rd + tfi * ki + EPSILON);
for (int s = 0; s < n_slots; s++) { int si = s * n_embd + i; float wg = lc->write_gates[t].data[s];
lc->num_states[t+1].data[si] = combined * lc->num_states[t].data[si] + wg * kv;
lc->den_states[t+1].data[si] = combined * lc->den_states[t].data[si] + wg * ki; }
}
}
}
}
for (int i = 0; i < lc->wkv.size; i++) lc->wkv_r.data[i] = lc->wkv.data[i] * lc->r.data[i];
matmul(&lc->tm_out, &lc->wkv_r, &lp->Wo);
vec_add(lc->x_after_tm.data, cache->x_final.data, lc->tm_out.data, cache->x_final.size);
layer_norm_seq(&lc->x_ln2, &lc->x_after_tm, &lp->ln2_weight, &lp->ln2_bias);
sigmoid_vec(lc->cm_mix_sig.data, lp->channel_mix.data, n_embd);
for (int t = 0; t < seq_len; t++) {
int i = 0;
__m256 vzero = _mm256_setzero_ps();
for (; i + 8 <= n_embd; i += 8) {
int idx = t * n_embd + i;
__m256 mix = _mm256_loadu_ps(lc->cm_mix_sig.data + i);
__m256 cur = _mm256_loadu_ps(lc->x_ln2.data + idx);
__m256 prev = (t > 0) ? _mm256_loadu_ps(lc->x_ln2.data + idx - n_embd) : vzero;
_mm256_storeu_ps(lc->xm.data + idx,
_mm256_fmadd_ps(cur, mix, _mm256_mul_ps(prev, _mm256_sub_ps(_mm256_set1_ps(1.0f), mix))));
}
for (; i < n_embd; i++) {
int idx = t * n_embd + i; float mix = lc->cm_mix_sig.data[i];
lc->xm.data[idx] = lc->x_ln2.data[idx] * mix + ((t > 0) ? lc->x_ln2.data[(t-1) * n_embd + i] : 0.0f) * (1.0f - mix);
}
}
{
hp_sgemm(lc->ffn_gu.data, lc->xm.data, lp->ffn_gate_up.data, seq_len, 2 * ffn_h, n_embd);
for (int t = 0; t < seq_len; t++) {
const float *gu = lc->ffn_gu.data + (size_t)t * 2 * ffn_h;
float *gp = lc->gate_pre.data + (size_t)t * ffn_h;
float *up = lc->up_val.data + (size_t)t * ffn_h;
float *gs = lc->gate_silu.data + (size_t)t * ffn_h;
float *hd = lc->hidden.data + (size_t)t * ffn_h;
int i = 0;
for (; i + 8 <= ffn_h; i += 8) {
__m256 x = _mm256_loadu_ps(gu + i);
__m256 u = _mm256_loadu_ps(gu + ffn_h + i);
__m256 s = _mm256_mul_ps(x, sigmoid256(x));
_mm256_storeu_ps(gs + i, s);
_mm256_storeu_ps(hd + i, _mm256_mul_ps(s, u));
_mm256_storeu_ps(gp + i, x);
_mm256_storeu_ps(up + i, u);
}
for (; i < ffn_h; i++) {
gs[i] = silu_f(gu[i]);
hd[i] = gs[i] * gu[ffn_h + i];
gp[i] = gu[i];
up[i] = gu[ffn_h + i];
}
}
}
matmul(&lc->cm_out, &lc->hidden, &lp->ffn_down);
vec_add(cache->x_final.data, lc->x_after_tm.data, lc->cm_out.data, cache->x_final.size);
}
layer_norm_seq(&cache->x_ln_out, &cache->x_final, &mp->ln_out_weight, &mp->ln_out_bias);
matmul(&cache->logits, &cache->x_ln_out, &mp->head);
return cross_entropy_loss(&cache->logits, tokens + 1, seq_len - 1);
}
static void backward_pass(ModelGrads *grads, const ForwardCache *cache,
const int *tokens, int seq_len,
const ModelParams *mp, const lrnnConfig *cfg) {
int n_embd = cfg->n_embd, vocab_size = cfg->vocab_size, ffn_h = ffn_hidden(cfg), lora_rank = cfg->decay_lora_rank;
int target_len = seq_len - 1;
bwd_reset();
Tensor d_logits = bwd_tensor(target_len, vocab_size);
Tensor d_x_ln_out = bwd_tensor(target_len, n_embd);
Tensor d_x_final = bwd_tensor(target_len, n_embd);
Tensor logits_view = { .data = cache->logits.data, .rows = target_len, .cols = vocab_size, .size = target_len * vocab_size };
softmax_cross_entropy_backward(&d_logits, &logits_view, tokens + 1);
matmul_backward_x(&d_x_ln_out, &d_logits, &mp->head);
Tensor x_ln_out_view = { .data = cache->x_ln_out.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
matmul_backward_w(&grads->head, &d_logits, &x_ln_out_view);
Tensor x_final_view = { .data = cache->x_final.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
layer_norm_backward_seq(&d_x_final, &grads->ln_out_weight, &grads->ln_out_bias, &d_x_ln_out, &x_final_view, &mp->ln_out_weight);
for (int layer_idx = mp->n_layers - 1; layer_idx >= 0; layer_idx--) {
const LayerParams *lp = &mp->layers[layer_idx]; const LayerCache *lc = &cache->layers[layer_idx]; LayerGrads *lg = &grads->layers[layer_idx];
Tensor d_cm_out = bwd_tensor(target_len, n_embd), d_hidden = bwd_tensor(target_len, ffn_h);
Tensor d_gate_silu = bwd_tensor(target_len, ffn_h);
Tensor d_xm = bwd_tensor(target_len, n_embd);
Tensor d_x_ln2 = bwd_tensor(target_len, n_embd), d_x_after_tm = bwd_tensor(target_len, n_embd);
tensor_copy(&d_cm_out, &d_x_final);
Tensor hidden_view = { .data = lc->hidden.data, .rows = target_len, .cols = ffn_h, .size = target_len * ffn_h };
matmul_backward_x(&d_hidden, &d_cm_out, &lp->ffn_down); matmul_backward_w(&lg->ffn_down, &d_cm_out, &hidden_view);
{
Tensor d_gu = bwd_tensor(target_len, 2 * ffn_h);
for (int i = 0; i < target_len * ffn_h; i++) {
d_gate_silu.data[i] = d_hidden.data[i] * lc->up_val.data[i];
}
for (int t = 0; t < target_len; t++) {
float *du = d_gu.data + (size_t)t * 2 * ffn_h + ffn_h;
const float *dh = d_hidden.data + (size_t)t * ffn_h;
const float *gs = lc->gate_silu.data + (size_t)t * ffn_h;
int j = 0;
for (; j + 8 <= ffn_h; j += 8) {
_mm256_storeu_ps(du + j, _mm256_mul_ps(_mm256_loadu_ps(dh + j), _mm256_loadu_ps(gs + j)));
}
for (; j < ffn_h; j++) du[j] = dh[j] * gs[j];
}
for (int t = 0; t < target_len; t++) {
float *dg = d_gu.data + (size_t)t * 2 * ffn_h;
const float *ds = d_gate_silu.data + (size_t)t * ffn_h;
const float *gp = lc->gate_pre.data + (size_t)t * ffn_h;
for (int j = 0; j < ffn_h; j++) {
float s = sigmoid_f(gp[j]);
dg[j] = ds[j] * s * (1.0f + gp[j] * (1.0f - s));
}
}
Tensor xm_view = { .data = lc->xm.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
matmul_backward_x(&d_xm, &d_gu, &lp->ffn_gate_up);
matmul_backward_w(&lg->ffn_gate_up, &d_gu, &xm_view);
BWD_FREE(&d_gu);
}
Tensor x_ln2_view = { .data = lc->x_ln2.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
channel_mix_shift_backward(&d_x_ln2, &lg->channel_mix, &d_xm, &x_ln2_view, &lc->cm_mix_sig, target_len, n_embd);
Tensor x_after_tm_view = { .data = lc->x_after_tm.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
layer_norm_backward_seq(&d_x_after_tm, &lg->ln2_weight, &lg->ln2_bias, &d_x_ln2, &x_after_tm_view, &lp->ln2_weight);
for (int i = 0; i < target_len * n_embd; i++) d_x_after_tm.data[i] += d_x_final.data[i];
BWD_FREE(&d_cm_out); BWD_FREE(&d_hidden); BWD_FREE(&d_gate_silu); BWD_FREE(&d_xm); BWD_FREE(&d_x_ln2);
Tensor d_tm_out = bwd_tensor(target_len, n_embd), d_wkv_r = bwd_tensor(target_len, n_embd);
Tensor d_wkv = bwd_tensor(target_len, n_embd), d_r = bwd_tensor(target_len, n_embd);
Tensor d_k_exp = bwd_tensor(target_len, n_embd), d_v = bwd_tensor(target_len, n_embd);
Tensor d_decay = bwd_tensor(target_len, n_embd), d_time_first_exp = tensor_alloc_1d(n_embd);
Tensor d_decay_pre = bwd_tensor(target_len, n_embd), d_decay_delta = bwd_tensor(target_len, n_embd);
Tensor d_decay_tmp = bwd_tensor(target_len, lora_rank), d_x_ln1_decay = bwd_tensor(target_len, n_embd);
Tensor d_r_pre = bwd_tensor(target_len, n_embd), d_k_pre = bwd_tensor(target_len, n_embd);
Tensor d_xr = bwd_tensor(target_len, n_embd), d_xk = bwd_tensor(target_len, n_embd), d_xv = bwd_tensor(target_len, n_embd);
Tensor d_x_ln1 = bwd_tensor(target_len, n_embd), d_x_shifted = bwd_tensor(target_len, n_embd);
Tensor d_x_layer_in = bwd_tensor(target_len, n_embd);
tensor_copy(&d_tm_out, &d_x_after_tm);
Tensor wkv_r_view = { .data = lc->wkv_r.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
matmul_backward_x(&d_wkv_r, &d_tm_out, &lp->Wo); matmul_backward_w(&lg->Wo, &d_tm_out, &wkv_r_view);
for (int i = 0; i < target_len * n_embd; i++) { d_wkv.data[i] = d_wkv_r.data[i] * lc->r.data[i]; d_r.data[i] = d_wkv_r.data[i] * lc->wkv.data[i]; }
Tensor x_ln1_view = { .data = lc->x_ln1.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
Tensor d_gate_write_logits = bwd_tensor(target_len, cfg->n_mem_slots);
Tensor d_gate_read_logits = bwd_tensor(target_len, cfg->n_mem_slots);
wkv_backward(&d_k_exp, &d_v, &d_decay, &d_time_first_exp, &d_gate_write_logits, &d_gate_read_logits,
&d_wkv, &lc->k_exp, &lc->v, &lc->decay, &lc->time_first_exp,
lc->num_states, lc->den_states, lc->write_gates, lc->read_gates,
target_len, n_embd, cfg->n_head, cfg->n_mem_slots, lp->alibi_slopes.data);
Tensor d_x_ln1_wgate = bwd_tensor(target_len, n_embd);
matmul_backward_x(&d_x_ln1_wgate, &d_gate_write_logits, &lp->mem_gate_write);
matmul_backward_w(&lg->mem_gate_write, &d_gate_write_logits, &x_ln1_view);
Tensor d_x_ln1_rgate = bwd_tensor(target_len, n_embd);
matmul_backward_x(&d_x_ln1_rgate, &d_gate_read_logits, &lp->mem_gate_read);
matmul_backward_w(&lg->mem_gate_read, &d_gate_read_logits, &x_ln1_view);
for (int i = 0; i < n_embd; i++) { float x = lp->time_first.data[i]; if (x >= -10.0f && x <= 10.0f) lg->time_first.data[i] += d_time_first_exp.data[i] * lc->time_first_exp.data[i]; }
for (int i = 0; i < target_len * n_embd; i++) { float y = lc->decay.data[i]; d_decay_pre.data[i] = d_decay.data[i] * y * (1.0f - y); }
for (int t = 0; t < target_len; t++) for (int i = 0; i < n_embd; i++) { int idx = t * n_embd + i; lg->decay_base.data[i] += d_decay_pre.data[idx]; d_decay_delta.data[idx] = d_decay_pre.data[idx]; }
Tensor decay_tmp_view = { .data = lc->decay_tmp.data, .rows = target_len, .cols = lora_rank, .size = target_len * lora_rank };
matmul_backward_x(&d_decay_tmp, &d_decay_delta, &lp->decay_lora_b); matmul_backward_w(&lg->decay_lora_b, &d_decay_delta, &decay_tmp_view);
matmul_backward_x(&d_x_ln1_decay, &d_decay_tmp, &lp->decay_lora_a); matmul_backward_w(&lg->decay_lora_a, &d_decay_tmp, &x_ln1_view);
sigmoid_backward(d_r_pre.data, d_r.data, lc->r.data, target_len * n_embd);
exp_backward_clamped(d_k_pre.data, d_k_exp.data, lc->k_pre.data, target_len * n_embd);
Tensor xr_view = { .data = lc->xr.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
Tensor xk_view = { .data = lc->xk.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
Tensor xv_view = { .data = lc->xv.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
matmul_backward_x(&d_xr, &d_r_pre, &lp->Wr); matmul_backward_w(&lg->Wr, &d_r_pre, &xr_view);
matmul_backward_x(&d_xk, &d_k_pre, &lp->Wk); matmul_backward_w(&lg->Wk, &d_k_pre, &xk_view);
matmul_backward_x(&d_xv, &d_v, &lp->Wv); matmul_backward_w(&lg->Wv, &d_v, &xv_view);
token_mixing_backward(&d_x_ln1, &d_x_shifted, &lg->time_mix_r, &lg->time_mix_k, &lg->time_mix_v,
&d_xr, &d_xk, &d_xv, &lc->x_ln1, &lc->x_shifted, &lc->mix_r_sig, &lc->mix_k_sig, &lc->mix_v_sig, target_len, n_embd);
for (int i = 0; i < target_len * n_embd; i++) d_x_ln1.data[i] += d_x_ln1_decay.data[i] + d_x_ln1_wgate.data[i] + d_x_ln1_rgate.data[i];
BWD_FREE(&d_gate_write_logits); BWD_FREE(&d_gate_read_logits); BWD_FREE(&d_x_ln1_wgate); BWD_FREE(&d_x_ln1_rgate);
Tensor d_x_ln1_shift = bwd_tensor(target_len, n_embd);
multi_scale_shift_backward(&d_x_ln1_shift, &lg->time_shift_w1, &lg->time_shift_w2, &lg->time_shift_w4,
&d_x_shifted, &lc->x_ln1, &lc->shift_w1_sig, &lc->shift_w2_sig, &lc->shift_w4_sig, &lc->shift_w_sum, target_len, n_embd);
for (int i = 0; i < target_len * n_embd; i++) d_x_ln1.data[i] += d_x_ln1_shift.data[i];
BWD_FREE(&d_x_ln1_shift);
Tensor layer_input;
if (layer_idx == 0) { layer_input.data = cache->x_ln0.data; layer_input.rows = target_len; layer_input.cols = n_embd; layer_input.size = target_len * n_embd; }
else { layer_input = bwd_tensor(target_len, n_embd); for (int i = 0; i < target_len * n_embd; i++) layer_input.data[i] = lc->x_after_tm.data[i] - lc->tm_out.data[i]; }
Tensor layer_input_view = { .data = lc->x_in.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
layer_norm_backward_seq(&d_x_layer_in, &lg->ln1_weight, &lg->ln1_bias, &d_x_ln1, &layer_input_view, &lp->ln1_weight);
if (layer_idx > 0) BWD_FREE(&layer_input);
for (int i = 0; i < target_len * n_embd; i++) d_x_layer_in.data[i] += d_x_after_tm.data[i];
tensor_copy(&d_x_final, &d_x_layer_in);
BWD_FREE(&d_tm_out); BWD_FREE(&d_wkv_r); BWD_FREE(&d_wkv); BWD_FREE(&d_r);
BWD_FREE(&d_k_exp); BWD_FREE(&d_v); BWD_FREE(&d_decay); BWD_FREE(&d_time_first_exp);
BWD_FREE(&d_decay_pre); BWD_FREE(&d_decay_delta); BWD_FREE(&d_decay_tmp); BWD_FREE(&d_x_ln1_decay);
BWD_FREE(&d_r_pre); BWD_FREE(&d_k_pre); BWD_FREE(&d_xr); BWD_FREE(&d_xk); BWD_FREE(&d_xv);
BWD_FREE(&d_x_ln1); BWD_FREE(&d_x_shifted); BWD_FREE(&d_x_layer_in); BWD_FREE(&d_x_after_tm);
}
Tensor d_emb_out = bwd_tensor(target_len, n_embd);
Tensor emb_out_view = { .data = cache->emb_out.data, .rows = target_len, .cols = n_embd, .size = target_len * n_embd };
layer_norm_backward_seq(&d_emb_out, &grads->ln0_weight, &grads->ln0_bias, &d_x_final, &emb_out_view, &mp->ln0_weight);
for (int t = 0; t < target_len; t++) { int tok = tokens[t]; for (int i = 0; i < n_embd; i++) grads->emb.data[tok * n_embd + i] += d_emb_out.data[t * n_embd + i]; }
BWD_FREE(&d_logits); BWD_FREE(&d_x_ln_out); BWD_FREE(&d_x_final); BWD_FREE(&d_emb_out);
}
typedef struct { float beta1, beta2, epsilon, weight_decay; int t; } AdamConfig;
typedef struct { Tensor m, v; } AdamState;
typedef struct {
AdamState emb, ln0_weight, ln0_bias;
struct { AdamState ln1_weight, ln1_bias, ln2_weight, ln2_bias;
AdamState time_shift_w1, time_shift_w2, time_shift_w4;
AdamState time_mix_r, time_mix_k, time_mix_v;
AdamState decay_lora_a, decay_lora_b, decay_base, time_first;
AdamState Wr, Wk, Wv, Wo, channel_mix;
AdamState ffn_gate_up, ffn_down;
AdamState mem_gate_write, mem_gate_read; } *layers;
AdamState ln_out_weight, ln_out_bias, head;
int n_layers;
} AdamStates;
static void init_adam_state(AdamState *as, int rows, int cols) { as->m = tensor_alloc(rows, cols); as->v = tensor_alloc(rows, cols); }
static void free_adam_state(AdamState *as) { tensor_free(&as->m); tensor_free(&as->v); }
static void init_adam_states(AdamStates *as, const lrnnConfig *cfg) {
int ne = cfg->n_embd, vs = cfg->vocab_size, fh = ffn_hidden(cfg), lr = cfg->decay_lora_rank;
as->n_layers = cfg->n_layer;
init_adam_state(&as->emb, vs, ne); init_adam_state(&as->ln0_weight, ne, 1); init_adam_state(&as->ln0_bias, ne, 1);
as->layers = calloc((size_t)cfg->n_layer, sizeof(*as->layers));
for (int i = 0; i < cfg->n_layer; i++) {
init_adam_state(&as->layers[i].ln1_weight, ne, 1); init_adam_state(&as->layers[i].ln1_bias, ne, 1);
init_adam_state(&as->layers[i].ln2_weight, ne, 1); init_adam_state(&as->layers[i].ln2_bias, ne, 1);
init_adam_state(&as->layers[i].time_shift_w1, ne, 1); init_adam_state(&as->layers[i].time_shift_w2, ne, 1); init_adam_state(&as->layers[i].time_shift_w4, ne, 1);
init_adam_state(&as->layers[i].time_mix_r, ne, 1); init_adam_state(&as->layers[i].time_mix_k, ne, 1); init_adam_state(&as->layers[i].time_mix_v, ne, 1);
init_adam_state(&as->layers[i].decay_lora_a, ne, lr); init_adam_state(&as->layers[i].decay_lora_b, lr, ne);
init_adam_state(&as->layers[i].decay_base, ne, 1); init_adam_state(&as->layers[i].time_first, ne, 1);
init_adam_state(&as->layers[i].Wr, ne, ne); init_adam_state(&as->layers[i].Wk, ne, ne);
init_adam_state(&as->layers[i].Wv, ne, ne); init_adam_state(&as->layers[i].Wo, ne, ne);
init_adam_state(&as->layers[i].channel_mix, ne, 1);
init_adam_state(&as->layers[i].ffn_gate_up, ne, 2 * fh); init_adam_state(&as->layers[i].ffn_down, fh, ne);
init_adam_state(&as->layers[i].mem_gate_write, ne, cfg->n_mem_slots); init_adam_state(&as->layers[i].mem_gate_read, ne, cfg->n_mem_slots);
}
init_adam_state(&as->ln_out_weight, ne, 1); init_adam_state(&as->ln_out_bias, ne, 1); init_adam_state(&as->head, ne, vs);
}
static void free_adam_states(AdamStates *as) {
free_adam_state(&as->emb); free_adam_state(&as->ln0_weight); free_adam_state(&as->ln0_bias);
for (int i = 0; i < as->n_layers; i++) {
free_adam_state(&as->layers[i].ln1_weight); free_adam_state(&as->layers[i].ln1_bias);
free_adam_state(&as->layers[i].ln2_weight); free_adam_state(&as->layers[i].ln2_bias);
free_adam_state(&as->layers[i].time_shift_w1); free_adam_state(&as->layers[i].time_shift_w2); free_adam_state(&as->layers[i].time_shift_w4);
free_adam_state(&as->layers[i].time_mix_r); free_adam_state(&as->layers[i].time_mix_k); free_adam_state(&as->layers[i].time_mix_v);
free_adam_state(&as->layers[i].decay_lora_a); free_adam_state(&as->layers[i].decay_lora_b);
free_adam_state(&as->layers[i].decay_base); free_adam_state(&as->layers[i].time_first);
free_adam_state(&as->layers[i].Wr); free_adam_state(&as->layers[i].Wk); free_adam_state(&as->layers[i].Wv); free_adam_state(&as->layers[i].Wo);
free_adam_state(&as->layers[i].channel_mix);
free_adam_state(&as->layers[i].ffn_gate_up); free_adam_state(&as->layers[i].ffn_down);
free_adam_state(&as->layers[i].mem_gate_write); free_adam_state(&as->layers[i].mem_gate_read);
}
free(as->layers); as->layers = NULL;
free_adam_state(&as->ln_out_weight); free_adam_state(&as->ln_out_bias); free_adam_state(&as->head);
}
static void adam_update(Tensor *param, Tensor *grad, AdamState *state, AdamConfig *config, float lr) {
float b1 = config->beta1, b2 = config->beta2, eps = config->epsilon, wd = config->weight_decay;
float bc1_inv = 1.0f / (1.0f - powf(b1, (float)config->t));
float bc2_inv = 1.0f / (1.0f - powf(b2, (float)config->t));
__m256 vb1 = _mm256_set1_ps(b1), vb2 = _mm256_set1_ps(b2);
__m256 v1b1 = _mm256_set1_ps(1.0f - b1), v1b2 = _mm256_set1_ps(1.0f - b2);
__m256 vbc1i = _mm256_set1_ps(bc1_inv), vbc2i = _mm256_set1_ps(bc2_inv);
__m256 veps = _mm256_set1_ps(eps), vlr = _mm256_set1_ps(lr), vwd = _mm256_set1_ps(wd);
int i = 0;
for (; i + 8 <= param->size; i += 8) {
__m256 g = _mm256_loadu_ps(grad->data + i);
__m256 m = _mm256_loadu_ps(state->m.data + i);
__m256 v = _mm256_loadu_ps(state->v.data + i);
__m256 p = _mm256_loadu_ps(param->data + i);
m = _mm256_add_ps(_mm256_mul_ps(vb1, m), _mm256_mul_ps(v1b1, g));
v = _mm256_add_ps(_mm256_mul_ps(vb2, v), _mm256_mul_ps(v1b2, _mm256_mul_ps(g, g)));
__m256 m_hat = _mm256_mul_ps(m, vbc1i), v_hat = _mm256_mul_ps(v, vbc2i);
__m256 upd = _mm256_div_ps(m_hat, _mm256_add_ps(_mm256_sqrt_ps(v_hat), veps));
p = _mm256_sub_ps(p, _mm256_mul_ps(vlr, _mm256_add_ps(upd, _mm256_mul_ps(vwd, p))));
_mm256_storeu_ps(state->m.data + i, m);
_mm256_storeu_ps(state->v.data + i, v);
_mm256_storeu_ps(param->data + i, p);
}
for (; i < param->size; i++) {
float g = grad->data[i];
state->m.data[i] = b1 * state->m.data[i] + (1.0f - b1) * g;
state->v.data[i] = b2 * state->v.data[i] + (1.0f - b2) * g * g;
float m_hat = state->m.data[i] * bc1_inv, v_hat = state->v.data[i] * bc2_inv;
param->data[i] -= lr * (m_hat / (sqrtf(v_hat) + eps) + wd * param->data[i]);
}
}
static float compute_tensor_norm_sq(const Tensor *t) { float sum = 0.0f; for (int i = 0; i < t->size; i++) sum += t->data[i] * t->data[i]; return sum; }
static void scale_tensor(Tensor *t, float scale) { for (int i = 0; i < t->size; i++) t->data[i] *= scale; }
static void clip_gradients_by_global_norm(ModelGrads *grads, float max_norm) {
float total_norm_sq = compute_tensor_norm_sq(&grads->emb) + compute_tensor_norm_sq(&grads->ln0_weight) + compute_tensor_norm_sq(&grads->ln0_bias);
for (int i = 0; i < grads->n_layers; i++) {
LayerGrads *lg = &grads->layers[i];
total_norm_sq += compute_tensor_norm_sq(&lg->ln1_weight) + compute_tensor_norm_sq(&lg->ln1_bias) + compute_tensor_norm_sq(&lg->ln2_weight) + compute_tensor_norm_sq(&lg->ln2_bias);
total_norm_sq += compute_tensor_norm_sq(&lg->time_shift_w1) + compute_tensor_norm_sq(&lg->time_shift_w2) + compute_tensor_norm_sq(&lg->time_shift_w4);
total_norm_sq += compute_tensor_norm_sq(&lg->time_mix_r) + compute_tensor_norm_sq(&lg->time_mix_k) + compute_tensor_norm_sq(&lg->time_mix_v);
total_norm_sq += compute_tensor_norm_sq(&lg->decay_lora_a) + compute_tensor_norm_sq(&lg->decay_lora_b) + compute_tensor_norm_sq(&lg->decay_base) + compute_tensor_norm_sq(&lg->time_first);
total_norm_sq += compute_tensor_norm_sq(&lg->Wr) + compute_tensor_norm_sq(&lg->Wk) + compute_tensor_norm_sq(&lg->Wv) + compute_tensor_norm_sq(&lg->Wo);
total_norm_sq += compute_tensor_norm_sq(&lg->channel_mix) + compute_tensor_norm_sq(&lg->ffn_gate_up) + compute_tensor_norm_sq(&lg->ffn_down);
total_norm_sq += compute_tensor_norm_sq(&lg->mem_gate_write) + compute_tensor_norm_sq(&lg->mem_gate_read);
}
total_norm_sq += compute_tensor_norm_sq(&grads->ln_out_weight) + compute_tensor_norm_sq(&grads->ln_out_bias) + compute_tensor_norm_sq(&grads->head);
float total_norm = sqrtf(total_norm_sq);
if (total_norm > max_norm) {
float scale = max_norm / (total_norm + 1e-8f);
scale_tensor(&grads->emb, scale); scale_tensor(&grads->ln0_weight, scale); scale_tensor(&grads->ln0_bias, scale);
for (int i = 0; i < grads->n_layers; i++) {
LayerGrads *lg = &grads->layers[i];
scale_tensor(&lg->ln1_weight, scale); scale_tensor(&lg->ln1_bias, scale); scale_tensor(&lg->ln2_weight, scale); scale_tensor(&lg->ln2_bias, scale);
scale_tensor(&lg->time_shift_w1, scale); scale_tensor(&lg->time_shift_w2, scale); scale_tensor(&lg->time_shift_w4, scale);
scale_tensor(&lg->time_mix_r, scale); scale_tensor(&lg->time_mix_k, scale); scale_tensor(&lg->time_mix_v, scale);
scale_tensor(&lg->decay_lora_a, scale); scale_tensor(&lg->decay_lora_b, scale); scale_tensor(&lg->decay_base, scale); scale_tensor(&lg->time_first, scale);
scale_tensor(&lg->Wr, scale); scale_tensor(&lg->Wk, scale); scale_tensor(&lg->Wv, scale); scale_tensor(&lg->Wo, scale);
scale_tensor(&lg->channel_mix, scale); scale_tensor(&lg->ffn_gate_up, scale); scale_tensor(&lg->ffn_down, scale);
scale_tensor(&lg->mem_gate_write, scale); scale_tensor(&lg->mem_gate_read, scale);
}
scale_tensor(&grads->ln_out_weight, scale); scale_tensor(&grads->ln_out_bias, scale); scale_tensor(&grads->head, scale);
}
}
static void apply_adam_updates(ModelParams *mp, ModelGrads *grads, AdamStates *adam, AdamConfig *config, float lr) {
config->t++;
adam_update(&mp->emb, &grads->emb, &adam->emb, config, lr);
adam_update(&mp->ln0_weight, &grads->ln0_weight, &adam->ln0_weight, config, lr);
adam_update(&mp->ln0_bias, &grads->ln0_bias, &adam->ln0_bias, config, lr);
for (int i = 0; i < mp->n_layers; i++) {
LayerParams *lp = &mp->layers[i]; LayerGrads *lg = &grads->layers[i];
adam_update(&lp->ln1_weight, &lg->ln1_weight, &adam->layers[i].ln1_weight, config, lr);
adam_update(&lp->ln1_bias, &lg->ln1_bias, &adam->layers[i].ln1_bias, config, lr);
adam_update(&lp->ln2_weight, &lg->ln2_weight, &adam->layers[i].ln2_weight, config, lr);
adam_update(&lp->ln2_bias, &lg->ln2_bias, &adam->layers[i].ln2_bias, config, lr);
adam_update(&lp->time_shift_w1, &lg->time_shift_w1, &adam->layers[i].time_shift_w1, config, lr);
adam_update(&lp->time_shift_w2, &lg->time_shift_w2, &adam->layers[i].time_shift_w2, config, lr);
adam_update(&lp->time_shift_w4, &lg->time_shift_w4, &adam->layers[i].time_shift_w4, config, lr);
adam_update(&lp->time_mix_r, &lg->time_mix_r, &adam->layers[i].time_mix_r, config, lr);
adam_update(&lp->time_mix_k, &lg->time_mix_k, &adam->layers[i].time_mix_k, config, lr);
adam_update(&lp->time_mix_v, &lg->time_mix_v, &adam->layers[i].time_mix_v, config, lr);
adam_update(&lp->decay_lora_a, &lg->decay_lora_a, &adam->layers[i].decay_lora_a, config, lr);
adam_update(&lp->decay_lora_b, &lg->decay_lora_b, &adam->layers[i].decay_lora_b, config, lr);
adam_update(&lp->decay_base, &lg->decay_base, &adam->layers[i].decay_base, config, lr);
adam_update(&lp->time_first, &lg->time_first, &adam->layers[i].time_first, config, lr);
adam_update(&lp->Wr, &lg->Wr, &adam->layers[i].Wr, config, lr); adam_update(&lp->Wk, &lg->Wk, &adam->layers[i].Wk, config, lr);
adam_update(&lp->Wv, &lg->Wv, &adam->layers[i].Wv, config, lr); adam_update(&lp->Wo, &lg->Wo, &adam->layers[i].Wo, config, lr);
adam_update(&lp->channel_mix, &lg->channel_mix, &adam->layers[i].channel_mix, config, lr);
adam_update(&lp->ffn_gate_up, &lg->ffn_gate_up, &adam->layers[i].ffn_gate_up, config, lr);
adam_update(&lp->ffn_down, &lg->ffn_down, &adam->layers[i].ffn_down, config, lr);
adam_update(&lp->mem_gate_write, &lg->mem_gate_write, &adam->layers[i].mem_gate_write, config, lr);
adam_update(&lp->mem_gate_read, &lg->mem_gate_read, &adam->layers[i].mem_gate_read, config, lr);
}
adam_update(&mp->ln_out_weight, &grads->ln_out_weight, &adam->ln_out_weight, config, lr);
adam_update(&mp->ln_out_bias, &grads->ln_out_bias, &adam->ln_out_bias, config, lr);
adam_update(&mp->head, &grads->head, &adam->head, config, lr);
}
static void train_model(const char *corpus_path, const char *save_path, int epochs, lrnnConfig *cfg, float lr, TokenizerType tok_type, bool auto_config) {
printf("======================================================================\n");
printf(" lrnn-like Model - Training (Enhanced C Implementation)\n");
printf(" High-perf GEMM: AVX2/FMA 16x6 packed microkernel, %d threads\n", g_pool_initialized ? g_pool_thread_count : (int)sysconf(_SC_NPROCESSORS_ONLN));
printf("======================================================================\n\n");
printf("Loading corpus: %s\n", corpus_path);
FILE *f = fopen(corpus_path, "rb");
if (!f) { fprintf(stderr, "Error: cannot open corpus file: %s\n", corpus_path); return; }
fseek(f, 0, SEEK_END); long file_size = ftell(f); fseek(f, 0, SEEK_SET);
if (file_size <= 0) { fprintf(stderr, "Error: empty or invalid corpus file\n"); fclose(f); return; }
char *text = (char *)malloc((size_t)file_size + 1);
size_t read_size = fread(text, 1, (size_t)file_size, f); text[read_size] = '\0'; fclose(f);
printf(" Loaded %zu bytes\n", read_size);
printf("\nBuilding tokenizer...\n");
Tokenizer tok; init_tokenizer(&tok, tok_type);
build_tokenizer(&tok, text, read_size, tok_type);
int vocab_size = tokenizer_vocab_size(&tok);
const char *tok_name = tok.type == TOKENIZER_CHAR ? "character" : (tok.type == TOKENIZER_WORD ? "word" : "BPE");
printf(" Vocabulary size: %d %s tokens\n", vocab_size, tok_name);
if (auto_config) { printf("\nAuto-configuring model for corpus size...\n"); *cfg = config_for_corpus(file_size, tok.type, vocab_size); }
else cfg->vocab_size = vocab_size;
int token_count; int *tokens = tokenizer_encode(&tok, text, read_size, &token_count); free(text);
printf(" Token count: %d\n", token_count);
if (tok.type == TOKENIZER_WORD || tok.type == TOKENIZER_BPE) printf(" Compression ratio: %.2fx\n", (float)read_size / (float)token_count);
pool_init();
printf("\nInitializing model...\n");
printf(" Layers: %d, Dim: %d, FFN: %d, Ctx: %d, LoRA: %d, Heads: %d\n", cfg->n_layer, cfg->n_embd, ffn_hidden(cfg), cfg->ctx_len, cfg->decay_lora_rank, cfg->n_head);
ModelParams mp; memset(&mp, 0, sizeof(mp)); init_model_params(&mp, cfg);
long total_params = mp.emb.size + mp.ln0_weight.size + mp.ln0_bias.size;
for (int i = 0; i < mp.n_layers; i++) { LayerParams *lp = &mp.layers[i]; total_params += lp->ln1_weight.size + lp->ln1_bias.size + lp->ln2_weight.size + lp->ln2_bias.size + lp->time_shift_w1.size + lp->time_shift_w2.size + lp->time_shift_w4.size + lp->time_mix_r.size + lp->time_mix_k.size + lp->time_mix_v.size + lp->decay_lora_a.size + lp->decay_lora_b.size + lp->decay_base.size + lp->time_first.size + lp->Wr.size + lp->Wk.size + lp->Wv.size + lp->Wo.size + lp->channel_mix.size + lp->ffn_gate_up.size + lp->ffn_down.size; }
total_params += mp.ln_out_weight.size + mp.ln_out_bias.size + mp.head.size;
printf(" Total parameters: %ld (%.2f MB)\n", total_params, (float)total_params * sizeof(float) / (1024.0f * 1024.0f));
ModelGrads grads; memset(&grads, 0, sizeof(grads)); init_model_grads(&grads, &mp, cfg);
AdamStates adam; memset(&adam, 0, sizeof(adam)); init_adam_states(&adam, cfg);
AdamConfig adam_cfg = { .beta1 = 0.9f, .beta2 = 0.999f, .epsilon = 1e-8f, .weight_decay = 0.0f, .t = 0 };
int batch_size = cfg->ctx_len; if (batch_size > token_count - 1) batch_size = token_count - 1;
int n_batches = (token_count - 1) / batch_size; if (n_batches < 1) n_batches = 1;
ForwardCache cache; memset(&cache, 0, sizeof(cache)); init_forward_cache(&cache, batch_size, cfg);
printf("\n======================================================================\n");
printf("Starting training...\n Tokenizer: %s | Batch: %d tokens | Batches/epoch: %d | Threads: %d\n", tok_name, batch_size, n_batches, g_pool_thread_count);
printf("======================================================================\n\n");
srand((unsigned)time(NULL) ^ (unsigned)getpid());
int total_steps = epochs * n_batches;
int warmup_steps = 50;
if (warmup_steps > total_steps / 4) warmup_steps = total_steps / 4;
if (warmup_steps < 1) warmup_steps = 1;
float lr_min = lr * 0.1f;
time_t start_time = time(NULL); float best_loss = FLT_MAX;
for (int epoch = 0; epoch < epochs; epoch++) {
float epoch_loss = 0.0f; int batch_count = 0;
int offset = rand() % (batch_size > 10 ? 10 : 1);
for (int batch = 0; batch < n_batches; batch++) {
int start = offset + batch * batch_size;
if (start + batch_size >= token_count) continue;
zero_model_grads(&grads);
float loss = forward_with_cache(&cache, tokens + start, batch_size, &mp, cfg);
if (!isfinite(loss)) { printf("Warning: NaN/Inf at epoch %d batch %d. Skipping.\n", epoch+1, batch); continue; }
backward_pass(&grads, &cache, tokens + start, batch_size, &mp, cfg);
clip_gradients_by_global_norm(&grads, 5.0f);
int step = epoch * n_batches + batch;
float step_lr;
if (step < warmup_steps) step_lr = lr * (float)(step + 1) / (float)warmup_steps;
else {
float progress = (float)(step - warmup_steps) / (float)(total_steps - warmup_steps);
step_lr = lr_min + 0.5f * (lr - lr_min) * (1.0f + cosf(3.14159265f * progress));
}
apply_adam_updates(&mp, &grads, &adam, &adam_cfg, step_lr);
epoch_loss += loss; batch_count++;
if ((batch + 1) % 10 == 0 || batch == n_batches - 1) { printf("\r Epoch %d/%d - Batch %d/%d - Loss: %.4f", epoch+1, epochs, batch+1, n_batches, batch_count > 0 ? epoch_loss / (float)batch_count : 0.0f); fflush(stdout); }
}
float avg_loss = (batch_count > 0) ? epoch_loss / (float)batch_count : 0.0f;
time_t elapsed = time(NULL) - start_time;
printf("\n Epoch %d/%d - Loss: %.4f - PPL: %.2f - Time: %02d:%02d:%02d\n", epoch+1, epochs, avg_loss, expf(avg_loss), (int)(elapsed/3600), (int)((elapsed%3600)/60), (int)(elapsed%60));
if (avg_loss < best_loss) { best_loss = avg_loss; printf(" ** New best loss! **\n"); }
if ((epoch + 1) % 5 == 0 || epoch == epochs - 1) { printf(" Saving: %s\n", save_path); save_model(save_path, &mp, cfg, &tok); }
}
printf("\nSaving final model to: %s\n", save_path); save_model(save_path, &mp, cfg, &tok);
free(tokens); free_forward_cache(&cache); free_model_grads(&grads); free_adam_states(&adam); free_model_params(&mp); free_tokenizer(&tok);
}
typedef struct { float prob; int index; } ProbIndex;
static int prob_index_cmp_desc(const void *a, const void *b) { float pa = ((const ProbIndex *)a)->prob, pb = ((const ProbIndex *)b)->prob; return (pa > pb) ? -1 : (pa < pb) ? 1 : 0; }
static int sample_top_p(const float *probs, int vocab_size, float top_p) {
ProbIndex *pi = (ProbIndex *)malloc((size_t)vocab_size * sizeof(ProbIndex));
if (!pi) { int best = 0; for (int i = 1; i < vocab_size; i++) if (probs[i] > probs[best]) best = i; return best; }
for (int i = 0; i < vocab_size; i++) { pi[i].prob = probs[i]; pi[i].index = i; }
qsort(pi, (size_t)vocab_size, sizeof(ProbIndex), prob_index_cmp_desc);
float cumsum = 0.0f; int cutoff = vocab_size;
for (int i = 0; i < vocab_size; i++) { cumsum += pi[i].prob; if (cumsum >= top_p) { cutoff = i + 1; break; } }
if (cutoff < 1) cutoff = 1;
float sum = 0.0f; for (int i = 0; i < cutoff; i++) sum += pi[i].prob;
float r = ((float)rand() / (float)RAND_MAX) * sum, running = 0.0f; int sampled = pi[0].index;
for (int i = 0; i < cutoff; i++) { running += pi[i].prob; if (running >= r) { sampled = pi[i].index; break; } }
free(pi); return sampled;
}
static void generate_text(const char *model_path, const char *seed_text, int n_tokens, float temperature, float top_p) {
printf("======================================================================\n Text Generation\n======================================================================\n\n");
ModelParams mp; lrnnConfig cfg; Tokenizer tok;
memset(&mp, 0, sizeof(mp)); memset(&tok, 0, sizeof(tok));
if (load_model(model_path, &mp, &cfg, &tok) != 0) { fprintf(stderr, "Failed to load model\n"); return; }
pool_init();
const char *tok_name = tok.type == TOKENIZER_CHAR ? "character" : (tok.type == TOKENIZER_WORD ? "word" : "BPE");
printf(" Model: %s | Tokenizer: %s | Vocab: %d | %dL/%dD | %d threads\n", model_path, tok_name, cfg.vocab_size, cfg.n_layer, cfg.n_embd, g_pool_thread_count);
ModelState state; init_model_state(&state, &cfg);
printf("\nSeed: \"%s\" | Generating %d tokens (temp=%.2f, top_p=%.2f)\n======================================================================\n%s", seed_text, n_tokens, temperature, top_p, seed_text);
fflush(stdout);
float *logits = (float *)malloc((size_t)cfg.vocab_size * sizeof(float));
float *probs = (float *)malloc((size_t)cfg.vocab_size * sizeof(float));
int seed_token_count; int *seed_tokens = tokenizer_encode(&tok, seed_text, strlen(seed_text), &seed_token_count);
int last_token = 0;
for (int i = 0; i < seed_token_count; i++) { forward_single(logits, seed_tokens[i], &mp, &state, &cfg); last_token = seed_tokens[i]; }
free(seed_tokens);
srand((unsigned int)time(NULL));
char decode_buf[1024];
for (int i = 0; i < n_tokens; i++) {
forward_single(logits, last_token, &mp, &state, &cfg);
if (temperature != 1.0f) for (int j = 0; j < cfg.vocab_size; j++) logits[j] /= temperature;
softmax_vec(probs, logits, cfg.vocab_size);
int next_token = sample_top_p(probs, cfg.vocab_size, top_p);
tokenizer_decode_token(&tok, next_token, decode_buf, sizeof(decode_buf));
printf("%s", decode_buf); fflush(stdout);
last_token = next_token;
}
printf("\n======================================================================\n");
free(logits); free(probs); free_model_state(&state); free_model_params(&mp); free_tokenizer(&tok);
}
static void print_usage(const char *prog) {
printf("Usage:\n Training:\n %s --train corpus.txt --save model.bin [options]\n\n Generation:\n %s --load model.bin --seed \"text\" [options]\n\n", prog, prog);
printf("Options:\n");
printf(" --train FILE Path to training corpus\n --save FILE Path to save model\n --load FILE Path to load model\n --seed TEXT Seed text for generation\n");
printf(" --epochs N Training epochs (default: 20)\n --tokens N Tokens to generate (default: 200)\n");
printf(" --layers N Number of layers\n --dim N Embedding dimension\n --heads N Number of heads\n --ctx N Max context length\n");
printf(" --lr FLOAT Learning rate (default: 0.0003)\n --temp FLOAT Temperature (default: 0.8)\n --top_p FLOAT Top-p sampling (default: 0.9)\n");
printf(" --tokenizer TYPE char, word, bpe, or auto (default: auto)\n");
printf(" --auto-config Auto-configure model size (default: on)\n --no-auto-config Disable auto-configuration\n --help Show this help\n");
}
static void set_flush_to_zero(void) {
_MM_SET_FLUSH_ZERO_MODE(_MM_FLUSH_ZERO_ON);
_MM_SET_DENORMALS_ZERO_MODE(_MM_DENORMALS_ZERO_ON);
}
int main(int argc, char *argv[]) {
set_flush_to_zero();
pool_init();
static struct option long_options[] = {
{"train", required_argument, 0, 't'}, {"save", required_argument, 0, 's'},
{"load", required_argument, 0, 'l'}, {"seed", required_argument, 0, 'S'},
{"epochs", required_argument, 0, 'e'}, {"tokens", required_argument, 0, 'n'},
{"layers", required_argument, 0, 'L'}, {"dim", required_argument, 0, 'd'},
{"heads", required_argument, 0, 'h'}, {"ctx", required_argument, 0, 'c'},
{"lr", required_argument, 0, 'r'}, {"temp", required_argument, 0, 'T'},
{"top_p", required_argument, 0, 'p'}, {"tokenizer", required_argument, 0, 'k'},
{"auto-config", no_argument, 0, 'A'}, {"no-auto-config", no_argument, 0, 'N'},
{"help", no_argument, 0, 'H'}, {0, 0, 0, 0}
};
char *train_path = NULL, *save_path = NULL, *load_path = NULL, *seed_text = NULL;
int epochs = 20, n_tokens = 200; float lr = 0.0003f, temperature = 0.8f, top_p = 0.9f;
TokenizerType tok_type = TOKENIZER_AUTO; bool auto_config = true, manual_config = false;
lrnnConfig cfg = default_config();
int opt, option_index = 0;
while ((opt = getopt_long(argc, argv, "t:s:l:S:e:n:L:d:h:c:r:T:p:k:ANH", long_options, &option_index)) != -1) {
switch (opt) {
case 't': train_path = optarg; break; case 's': save_path = optarg; break;
case 'l': load_path = optarg; break; case 'S': seed_text = optarg; break;
case 'e': epochs = atoi(optarg); break; case 'n': n_tokens = atoi(optarg); break;
case 'L': cfg.n_layer = atoi(optarg); manual_config = true; break;
case 'd': cfg.n_embd = atoi(optarg); manual_config = true; break;
case 'h': cfg.n_head = atoi(optarg); manual_config = true; break;
case 'c': cfg.ctx_len = atoi(optarg); manual_config = true; break;
case 'r': lr = (float)atof(optarg); break; case 'T': temperature = (float)atof(optarg); break;
case 'p': top_p = (float)atof(optarg); break;
case 'k':
if (strcmp(optarg, "char") == 0) tok_type = TOKENIZER_CHAR;
else if (strcmp(optarg, "word") == 0) tok_type = TOKENIZER_WORD;
else if (strcmp(optarg, "bpe") == 0) tok_type = TOKENIZER_BPE;
else tok_type = TOKENIZER_AUTO;
break;
case 'A': auto_config = true; break; case 'N': auto_config = false; break;
case 'H': default: print_usage(argv[0]); return (opt == 'H') ? 0 : 1;
}
}
if (manual_config) auto_config = false;
if (train_path) {
if (!save_path) { fprintf(stderr, "Error: --save is required when training\n"); return 1; }
train_model(train_path, save_path, epochs, &cfg, lr, tok_type, auto_config);
} else if (load_path) {
if (!seed_text || strlen(seed_text) == 0) { fprintf(stderr, "Error: --seed is required for generation\n"); return 1; }
generate_text(load_path, seed_text, n_tokens, temperature, top_p);
} else { print_usage(argv[0]); return 1; }
return 0;
}