File size: 3,608 Bytes
9b9eafc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
// Nucleus (top-p) sampling chạy trực tiếp trên logits, không cần sắp xếp toàn bộ từ vựng,
// cùng các tiện ích đo tốc độ token.

export interface LogitsLike {
  dims: readonly number[];
  data: ArrayLike<number> & { subarray(begin: number, end: number): any; length: number };
}

/**
 * Trả về logits processor giữ lại tập token nhỏ nhất có tổng xác suất >= topP
 * (các token còn lại bị đặt -Infinity). Chỉnh sửa logits tại chỗ.
 */
export function createTopPProcessor(topP = 0.95) {
  if (!(topP > 0 && topP <= 1)) throw new Error('topP must be in (0, 1].');
  let scratch: Float32Array | undefined;

  return (_inputIds: unknown, logits: LogitsLike) => {
    if (topP === 1) return logits;
    const vocab = logits.dims[logits.dims.length - 1];
    if (!scratch || scratch.length !== vocab) scratch = new Float32Array(vocab);
    const t = scratch;

    for (let offset = 0; offset < logits.data.length; offset += vocab) {
      const row = logits.data.subarray(offset, offset + vocab) as Float32Array;

      let max = -Infinity;
      for (let i = 0; i < vocab; i++) max = Math.max(max, row[i]);
      if (!Number.isFinite(max)) throw new Error('The model produced invalid sampling scores.');

      let total = 0;
      for (let i = 0; i < vocab; i++) total += Math.exp(row[i] - max);
      const allowedTail = (1 - topP) * total; // khối lượng xác suất được phép loại bỏ

      // Tìm ngưỡng thô: loại các logit nhỏ hơn `threshold` mà tổng khối lượng vẫn <= allowedTail.
      let threshold = max - 8;
      let count = 0;
      let dropped = 0;
      for (;;) {
        count = 0;
        dropped = 0;
        for (let i = 0; i < vocab; i++) {
          if (row[i] >= threshold) t[count++] = row[i];
          else dropped += Math.exp(row[i] - max);
        }
        if (dropped <= allowedTail) break;
        threshold -= 8;
      }

      // Tinh chỉnh: loại thêm các ứng viên nhỏ nhất chừng nào khối lượng loại bỏ còn <= allowedTail.
      t.subarray(0, count).sort();
      let f = 0;
      while (f < count - 1) {
        const next = dropped + Math.exp(t[f] - max);
        if (next > allowedTail) break;
        dropped = next;
        f++;
      }

      // Xử lý các giá trị bằng nhau ở biên để đúng số lượng token bị loại.
      const cutoff = f ? t[f - 1] : threshold;
      let ties = 0;
      for (let i = f - 1; i >= 0 && t[i] === cutoff; i--) ties++;
      for (let i = 0; i < vocab; i++) {
        if (row[i] < cutoff || (row[i] === cutoff && ties-- > 0)) row[i] = -Infinity;
      }
    }
    return logits;
  };
}

/** Đo tốc độ token/giây theo cửa sổ trượt. */
export class RateMeter {
  total = 0;
  private firstAt: number | null = null;
  private points: Array<{ at: number; total: number }> = [];

  constructor(
    private windowMs = 1000,
    private minimumMs = 250,
  ) {}

  add(count: number, at: number) {
    if (count <= 0) return;
    this.firstAt ??= at;
    this.total += count;
    this.points.push({ at, total: this.total });
    this.prune(at);
  }

  private prune(now: number) {
    const cutoff = now - this.windowMs;
    while (this.points.length > 1 && this.points[1].at <= cutoff) this.points.shift();
  }

  rate(now: number): number | null {
    if (this.firstAt === null || now - this.firstAt < this.minimumMs) return null;
    this.prune(now);
    const span = Math.min(this.windowMs, now - this.firstAt);
    return ((this.total - this.points[0].total) * 1000) / span;
  }
}