File size: 7,687 Bytes
30bafb7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
// Model sharding — the piece DaisyChain-Train does not have.
//
// DaisyChain-Train pools COMPUTE: every node holds a full replica, so a model
// bigger than one machine cannot be trained. Inference is where that limit can
// be lifted, because a forward pass is a chain: layer l needs layer l-1's
// OUTPUT, never its WEIGHTS. So the layers live on different machines and the
// activation travels instead.
//
// And because safetensors gives every tensor an exact byte range, each device
// fetches only its own layers straight from the Hub. No device — not even the
// head — ever holds the whole model. That is what makes the pooling real
// rather than a redistribution of something one machine already had to load.
//
// The ring shape is forced by weight tying. When lm_head is tied to the
// embedding table, the largest tensor in the model is needed at BOTH ends: to
// embed the prompt and to produce the logits. Copying it to the last stage
// would hand back most of the memory just pooled, so the last stage returns
// its hidden state to the head, which owns the embedding once and does both
// ends of the pass.
(function (root) {
  "use strict";

  // ---- the plan ---------------------------------------------------------------
  // Stage 0 is the HEAD: embeddings, final norm, lm_head. Every stage may also
  // own a contiguous run of layers, apportioned by MEASURED capacity — the
  // same self-calibrating idea as cluster.py's capacity_score, except what is
  // balanced is layers per device rather than batch per device.
  //
  // Largest-remainder apportionment, so the plan is a pure function of the
  // capacity report: every device derives the same plan from the same inputs
  // and the plan never has to be trusted, only compared.
  function planStages(spec, caps) {
    const n = caps.length;
    if (n < 1) throw new Error("no devices to plan across");
    const total = caps.reduce((a, d) => a + Math.max(1e-6, d.capacity), 0);
    const exact = caps.map(d => spec.layers * Math.max(1e-6, d.capacity) / total);
    const floor = exact.map(Math.floor);
    let left = spec.layers - floor.reduce((a, b) => a + b, 0);
    const order = exact.map((e, i) => [e - floor[i], i]).sort((a, b) => b[0] - a[0] || a[1] - b[1]);
    for (let i = 0; i < order.length && left > 0; i++, left--) floor[order[i][1]]++;
    const stages = [];
    let lo = 0;
    for (let i = 0; i < n; i++) {
      const hi = lo + floor[i];
      stages.push({ id: caps[i].id, index: i, lo, hi, head: i === 0,
                    backend: caps[i].backend, capacity: caps[i].capacity });
      lo = hi;
    }
    // A device with no layers is only worth a hop if it is the head, which has
    // real work (embed + unembed). Anyone else empty is dropped.
    return stages.filter(s => s.head || s.hi > s.lo).map((s, i) => ({ ...s, index: i }));
  }

  // What a stage will cost to hold, computed from the file's own header BEFORE
  // any weight bytes move — so a device can see whether its slice fits before
  // spending the bandwidth finding out.
  function stageBytes(spec, available, st, ArchMod) {
    const names = ArchMod.tensorsFor(spec, available, st);
    let n = 0;
    for (const name of names) n += available.get(name).elems * 4;   // f32 in memory
    return n;
  }

  // ---- arranging fetched tensors into what the forward pass wants -------------
  // Two shape conventions have to be reconciled here, and getting it wrong is
  // silent: torch.nn.Linear stores weights (out, in) while the GEMM wants
  // (in, out), whereas GPT-2's Conv1D already stores (in, out). A wrong
  // transpose does not throw — it produces a model that generates confident
  // nonsense — so the layout comes from the spec and is applied once, at load.
  function toKN(w, outDim, inDim, layout) {
    if (layout === "in_out") return w;                    // already k x n
    const out = new Float32Array(w.length);
    for (let o = 0; o < outDim; o++)
      for (let i = 0; i < inDim; i++) out[i * outDim + o] = w[o * inDim + i];
    return out;
  }

  function stageWeights(spec, st, got, ArchMod) {
    const resolve = ArchMod.resolver(new Map([...got.keys()].map(k => [k, true])));
    const pick = (name, opt) => { const r = resolve(name, opt); return r ? got.get(r) : null; };
    const C = spec.hidden, L = spec.weightLayout;
    const w = { layers: [] };

    if (st.head) {
      const h = ArchMod.headTensors(spec);
      w.emb = pick(h.emb);                                // (vocab, C) — used as a lookup, not a GEMM
      w.pos = pick(h.pos, true);
      w.nrmF = pick(h.nrmF, true);
      w.nrmFb = pick(h.nrmFb, true);
      // lm_head as k x n = (C, vocab). When tied, that is embᵀ.
      const lm = pick(h.lmHead, true);
      w.lmHead = lm ? toKN(lm, spec.vocab, C, L) : transposeEmb(w.emb, spec.vocab, C);
      if (!w.nrmF) w.nrmF = new Float32Array(C).fill(1);  // a model without a final norm weight
    }

    const qDim = spec.heads * spec.headDim, kvDim = spec.kvHeads * spec.headDim;
    for (let l = st.lo; l < st.hi; l++) {
      const t = ArchMod.layerTensors(spec, l);
      const ly = { nrm1: pick(t.nrm1), nrm1b: pick(t.nrm1b, true),
                   nrm2: pick(t.nrm2), nrm2b: pick(t.nrm2b, true) };
      if (spec.qkvFused) {
        ly.Wqkv = toKN(pick(t.Wqkv), 3 * C, C, L);
        ly.bqkv = pick(t.bqkv, true);
      } else {
        ly.Wq = toKN(pick(t.Wq), qDim, C, L); ly.bq = pick(t.bq, true);
        ly.Wk = toKN(pick(t.Wk), kvDim, C, L); ly.bk = pick(t.bk, true);
        ly.Wv = toKN(pick(t.Wv), kvDim, C, L); ly.bv = pick(t.bv, true);
      }
      ly.Wo = toKN(pick(t.Wo), C, qDim, L); ly.bo = pick(t.bo, true);
      if (spec.gated) {
        ly.Wgate = toKN(pick(t.Wgate), spec.inter, C, L);
        ly.Wup = toKN(pick(t.Wup), spec.inter, C, L);
      } else {
        ly.Wfc = toKN(pick(t.Wfc), spec.inter, C, L); ly.bfc = pick(t.bfc, true);
      }
      ly.Wdown = toKN(pick(t.Wdown), C, spec.inter, L); ly.bdown = pick(t.bdown, true);
      w.layers.push(ly);
    }
    return w;
  }
  function transposeEmb(emb, vocab, C) {
    const out = new Float32Array(emb.length);
    for (let v = 0; v < vocab; v++)
      for (let c = 0; c < C; c++) out[c * vocab + v] = emb[v * C + c];
    return out;
  }

  // ---- hashes -----------------------------------------------------------------
  // FNV-1a over raw bytes, identical to the function DaisyChain-Web hashes
  // replicas with, so hashes stay comparable across the projects.
  function fnv1a(bytes) {
    let h = 0x811c9dc5;
    for (let i = 0; i < bytes.length; i++) { h ^= bytes[i]; h = Math.imul(h, 0x01000193); }
    return h >>> 0;
  }
  function hashF32(a) { return fnv1a(new Uint8Array(a.buffer, a.byteOffset, a.byteLength)); }

  // The model fingerprint every stage repeats in its status. Derived from the
  // repo id, revision and the tensor index — NOT from the weights, because no
  // device reads all of them. It answers "are we all running the same model?",
  // which is the question that matters when stages hold disjoint pieces.
  function modelFingerprint(repo, revision, tensors) {
    const parts = [repo, revision || "main"];
    for (const name of [...tensors.keys()].sort()) {
      const t = tensors.get(name);
      parts.push(`${name}:${t.dtype}:${t.shape.join("x")}:${t.start}:${t.end}`);
    }
    return fnv1a(new TextEncoder().encode(parts.join("|")));
  }

  const api = { planStages, stageBytes, stageWeights, toKN, transposeEmb,
                fnv1a, hashF32, modelFingerprint };
  if (typeof module !== "undefined" && module.exports) module.exports = api;
  else root.Shard = api;
})(typeof self !== "undefined" ? self : this);