File size: 4,882 Bytes
5fd9a42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0b29713
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// The chat loop, ported line-for-line from OnnxChat in
// scripts/quantize_web.py (the parity-proven python reference). No imports:
// the ONNX runtime, sessions, and tokenizer are injected so the same module
// runs in the browser and under Node for the parity test.

export const PAD_ID = 0;
export const BOS_ID = 1;
export const EOS_ID = 2;

const D = 384;
const K = 8; // thoughts per turn

export class OnnxChat {
  /**
   * @param ort onnxruntime-web module (for Tensor construction)
   * @param sessions {encoder, thinker, decoder} InferenceSessions
   * @param tok {encode(text)->number[], decode(ids)->string}
   */
  constructor(ort, sessions, tok) {
    this.ort = ort;
    this.s = sessions;
    this.tok = tok;
    this.history = [];
  }

  reset() {
    this.history = [];
  }

  ids64(ids) {
    return new this.ort.Tensor(
      "int64", BigInt64Array.from(ids, BigInt), [1, ids.length]);
  }

  /** Greedy reply; onToken(text-so-far) fires as tokens decode. */
  async reply(text, onToken) {
    this.history.push(text.trim());
    const turns = this.history.slice(-6);
    const n = turns.length;
    const firstRole = (this.history.length - n) % 2;
    const th = new Float32Array(n * K * D);
    const roles = [], dist = [];
    for (let j = 0; j < n; j++) {
      const ids = [BOS_ID, ...this.tok.encode(turns[j]).slice(0, 254), EOS_ID];
      const enc = await this.s.encoder.run({ ids: this.ids64(ids) });
      th.set(enc.thoughts.data, j * K * D);
      roles.push((firstRole + j) % 2);
      dist.push(Math.min(n - j, 6));
    }
    const out = await this.s.thinker.run({
      ctx_th: new this.ort.Tensor("float32", th, [1, n, K, D]),
      ctx_roles: this.ids64(roles),
      dist: this.ids64(dist),
    });
    const score = out.score.data;
    let best = 0;
    for (let h = 1; h < score.length; h++) if (score[h] < score[best]) best = h;
    const thoughts = new this.ort.Tensor(
      "float32", out.hyps.data.slice(best * K * D, (best + 1) * K * D),
      [1, K, D]);

    const ids = [BOS_ID];
    for (let step = 0; step < 255; step++) {
      const fed = ids.length < 2 ? [...ids, 0] : ids;
      const dec = await this.s.decoder.run({
        thoughts,
        ids: this.ids64(fed),
        pos: new this.ort.Tensor(
          "int64", BigInt64Array.from([BigInt(ids.length - 1)]), [1]),
      });
      const lg = dec.logits.data;
      if (ids.length >= 3) { // no_repeat_ngram=3
        const p0 = ids[ids.length - 2], p1 = ids[ids.length - 1];
        for (let k = 0; k < ids.length - 2; k++) {
          if (ids[k] === p0 && ids[k + 1] === p1) lg[ids[k + 2]] = -Infinity;
        }
      }
      let nxt = 0;
      for (let v = 1; v < lg.length; v++) if (lg[v] > lg[nxt]) nxt = v;
      ids.push(nxt);
      if (onToken) onToken(this.tok.decode(ids));
      if (nxt === EOS_ID) break;
    }
    const reply = this.tok.decode(ids);
    this.history.push(reply);
    return reply;
  }
}

// The paper's §6.5 matched token-LM baseline. Ported line-for-line from
// OnnxLmChat in scripts/quantize_web.py. Deliberately un-improved: flat
// token history (no turn windowing) and no repeat-ngram ban, because its
// repetition loops and apology-default register are the paper's documented
// finding about this paradigm, not bugs for this loop to paper over.
export class OnnxLmChat {
  /**
   * @param ort onnxruntime-web module
   * @param session {lm} InferenceSession
   * @param tok {encode(text)->number[], decode(ids)->string}
   * @param maxLen model's max_seq_len (384 for the released baseline)
   * @param maxNew max new tokens per reply (64 for the released baseline)
   */
  constructor(ort, session, tok, maxLen, maxNew = 64) {
    this.ort = ort;
    this.s = session;
    this.tok = tok;
    this.maxLen = maxLen;
    this.maxNew = maxNew;
    this.history = [];
  }

  reset() {
    this.history = [];
  }

  ids64(ids) {
    return new this.ort.Tensor(
      "int64", BigInt64Array.from(ids, BigInt), [1, ids.length]);
  }

  async reply(text, onToken) {
    this.history.push(text.trim());
    let ids = [];
    for (const t of this.history) {
      ids = ids.concat([BOS_ID], this.tok.encode(t), [EOS_ID]);
    }
    const room = this.maxLen - this.maxNew - 1;
    const out = ids.slice(-room);
    out.push(BOS_ID);
    const gen = [];
    for (let step = 0; step < this.maxNew; step++) {
      const dec = await this.s.lm.run({
        ids: this.ids64(out),
        pos: new this.ort.Tensor(
          "int64", BigInt64Array.from([BigInt(out.length - 1)]), [1]),
      });
      const lg = dec.logits.data;
      let nxt = 0;
      for (let v = 1; v < lg.length; v++) if (lg[v] > lg[nxt]) nxt = v;
      if (nxt === EOS_ID) break;
      gen.push(nxt);
      out.push(nxt);
      if (onToken) onToken(this.tok.decode(gen));
    }
    const reply = this.tok.decode(gen);
    this.history.push(reply);
    return reply;
  }
}