File size: 6,204 Bytes
c971a45 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | // Single source of numeric truth for the ACTIVE Marian model — shared by JS
// reference math, engine code, and templated into WGSL kernels.
//
// The geometry lives in `let` bindings booted to MoxhiMT-30 and swapped by
// applyModelConfig(manifest.model) when weights load (uploadParsed calls it),
// so ONE build serves the whole MT-30/60 family. Two rules keep that sound:
// 1. Consumers read these at CALL time (function bodies / default params) —
// never snapshot them into module-top-level derived consts.
// 2. Every dim baked into WGSL rides the pipeline-cache key (defines in
// normalizeFlags), so switching models builds fresh pipelines instead of
// reusing stale ones. runEncoder/createDecodeState assert the weights
// they're handed match the active config (assertModelActive) — using two
// models interleaved without re-applying is a caller error, caught loud.
export let D_MODEL, HEADS, HEAD_DIM, FFN, VOCAB;
export let ENC_LAYERS, DEC_LAYERS;
export let MAX_POS, SRC_CAP, DECODE_CAP;
// Attention kernel shared-memory scores capacity — max valid K/V length any
// dispatch may see (bucketed SRC_CAP and DECODE_CAP both fit), +32 headroom.
export let SCORES_CAP;
export let EOS, PAD, DECODER_START;
export let LN_EPS;
export let EMBED_SCALE; // √D_MODEL (scale_embedding), 1 when scaleEmbedding: false
export let ATTN_SCALE; // 1/√HEAD_DIM
export let BITMASK_WORDS; // ⌈VOCAB/32⌉ u32 words per repetition-bitmask row
// Decode policy (HF generation config), not model geometry — constant across
// the family so far; move into the manifest if a member ever changes it.
export const REP_PENALTY = 1.2;
// Validate a manifest `model` block and return it normalized. Throws on
// missing/ill-typed fields and on anything the kernels cannot serve — the
// engine's portability envelope is Marian with vec4-able dims, fused-QKV
// strides (heads·headDim == dModel), swish FFN, and tied embeddings (the
// lm_head reads shared.weight); those last two are the exporter's to assert.
const INT_FIELDS = [
'dModel', 'heads', 'headDim', 'ffn', 'encLayers', 'decLayers',
'vocab', 'maxPos', 'srcCap', 'decodeCap', 'eos', 'pad', 'decoderStart',
];
export function parseModelConfig(model) {
if (!model || typeof model !== 'object') throw new Error('model config: not an object');
const m = {};
for (const f of INT_FIELDS) {
const v = model[f];
if (!Number.isInteger(v) || v < 0) {
throw new Error(`model config: ${f} must be a non-negative integer, got ${v}`);
}
m[f] = v;
}
for (const f of INT_FIELDS.slice(0, 10)) { // all but eos/pad/decoderStart
if (m[f] === 0) throw new Error(`model config: ${f} must be positive`);
}
if (!(typeof model.lnEps === 'number' && model.lnEps > 0 && model.lnEps < 0.1)) {
throw new Error(`model config: lnEps must be in (0, 0.1), got ${model.lnEps}`);
}
m.lnEps = model.lnEps;
if (m.heads * m.headDim !== m.dModel) {
throw new Error(`model config: heads·headDim = ${m.heads * m.headDim} != dModel ${m.dModel}`);
}
for (const f of ['dModel', 'headDim', 'ffn']) {
if (m[f] % 4 !== 0) {
throw new Error(`model config: ${f} = ${m[f]} must be a multiple of 4 (vec4 kernels)`);
}
}
if (m.srcCap % 32 !== 0) {
throw new Error(`model config: srcCap ${m.srcCap} must be a multiple of 32 (bucketFor invariant)`);
}
if (m.srcCap > m.maxPos || m.decodeCap > m.maxPos) {
throw new Error(
`model config: srcCap ${m.srcCap} / decodeCap ${m.decodeCap} must not exceed maxPos ${m.maxPos}`);
}
for (const f of ['eos', 'pad', 'decoderStart']) {
if (m[f] >= m.vocab) throw new Error(`model config: ${f} = ${m[f]} out of vocab ${m.vocab}`);
}
// scale_embedding: false is representable (embed kernel takes EMBED_SCALE).
if (model.scaleEmbedding === false) m.scaleEmbedding = false;
m.embedScale = m.scaleEmbedding === false ? 1 : Math.sqrt(m.dModel);
return m;
}
let ACTIVE = null;
// The currently applied config, as a re-appliable copy — for policy caching
// keys and for tests that switch models and must restore what they found.
export function activeModelConfig() {
return { ...ACTIVE };
}
// Make `model` the active config. Idempotent; callers normally never invoke
// this directly — uploadParsed/loadWeights apply the manifest's block.
export function applyModelConfig(model) {
const m = parseModelConfig(model);
ACTIVE = m;
D_MODEL = m.dModel; HEADS = m.heads; HEAD_DIM = m.headDim; FFN = m.ffn; VOCAB = m.vocab;
ENC_LAYERS = m.encLayers; DEC_LAYERS = m.decLayers;
MAX_POS = m.maxPos; SRC_CAP = m.srcCap; DECODE_CAP = m.decodeCap;
SCORES_CAP = Math.ceil(Math.max(m.srcCap, m.decodeCap) / 32) * 32 + 32;
EOS = m.eos; PAD = m.pad; DECODER_START = m.decoderStart;
LN_EPS = m.lnEps;
EMBED_SCALE = m.embedScale;
ATTN_SCALE = 1 / Math.sqrt(m.headDim);
BITMASK_WORDS = Math.ceil(m.vocab / 32);
return m;
}
// Guard for engine entry points: the weights being dispatched must have been
// the last applied config (their dims are live in kernel defines right now).
export function assertModelActive(model, what = 'weights') {
const m = model ?? {};
const stale =
m.dModel !== D_MODEL || m.heads !== HEADS || m.headDim !== HEAD_DIM ||
m.ffn !== FFN || m.vocab !== VOCAB ||
m.encLayers !== ENC_LAYERS || m.decLayers !== DEC_LAYERS ||
m.maxPos !== MAX_POS || m.srcCap !== SRC_CAP || m.decodeCap !== DECODE_CAP ||
m.eos !== EOS || m.pad !== PAD || m.decoderStart !== DECODER_START;
if (stale) {
throw new Error(
`${what} belong to a different model config than the active one — ` +
'call applyModelConfig(weights.model) (loadWeights/uploadParsed do it) before dispatching');
}
}
// Boot config: MoxhiMT-30 zh→vi (d_model 448, 8 heads × 56, ffn 1792, enc 8 /
// dec 2, vocab 24k) — kept as the pre-load default so tools and unit tests
// see the shipped numbers (EMBED_SCALE √448, ATTN_SCALE 1/√56, SCORES_CAP 352).
applyModelConfig({
dModel: 448, heads: 8, headDim: 56, ffn: 1792,
encLayers: 8, decLayers: 2,
vocab: 24000, maxPos: 512, srcCap: 320, decodeCap: 224,
eos: 2, pad: 0, decoderStart: 0, lnEps: 1e-5,
});
|