File size: 5,111 Bytes
cd8bd0a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
/**
 * Transformers.js local embedding (D8) — Xenova/all-MiniLM-L6-v2.
 *
 * IMPORTANT: @huggingface/transformers is imported lazily (await import())
 * ONLY when this function is called. Never imported at module level.
 * This satisfies D8 + D25 (serverExternalPackages + no bundle impact).
 */

import { sanitizeErrorMessage } from "@omniroute/open-sse/utils/error.ts";
import type { EmbeddingResult, EmbeddingError } from "./types";

const TRANSFORMERS_MODEL =
  process.env.MEMORY_TRANSFORMERS_MODEL || "Xenova/all-MiniLM-L6-v2";

// Singleton pipeline, initialized once
type PipelineFn = (text: string | string[], options?: Record<string, unknown>) => Promise<unknown>;
let _pipeline: PipelineFn | null = null;
let _pipelineLoading: Promise<PipelineFn> | null = null;

/** For testing: inject a mock pipeline factory. */
export function _injectPipeline(fn: PipelineFn | null): void {
  _pipeline = fn;
  _pipelineLoading = null;
}

async function getOrLoadPipeline(): Promise<PipelineFn> {
  if (_pipeline) return _pipeline;
  if (_pipelineLoading) return _pipelineLoading;

  _pipelineLoading = (async (): Promise<PipelineFn> => {
    // Lazy import — never at module level (D8, D25)
    const transformers = await import("@huggingface/transformers");
    const { pipeline } = transformers as { pipeline: (task: string, model: string, opts?: Record<string, unknown>) => Promise<PipelineFn> };
    const pipe = await pipeline("feature-extraction", TRANSFORMERS_MODEL, { dtype: "q8" });
    _pipeline = pipe;
    _pipelineLoading = null;
    return pipe;
  })();

  return _pipelineLoading;
}

/**
 * Convert Tensor-like output from transformers pipeline to Float32Array.
 * Transformers.js pipelines return a Tensor with `.data` (Float32Array or similar)
 * and `.dims` [batch, seq, hidden_size]. We flatten to hidden_size via mean pooling.
 */
function tensorToFloat32Array(output: unknown): Float32Array {
  // Handle Tensor objects from @huggingface/transformers
  const tensor = output as {
    data?: Float32Array | number[];
    dims?: number[];
    tolist?: () => number[][][];
  };

  if (tensor && tensor.data && tensor.dims) {
    const data = tensor.data instanceof Float32Array ? tensor.data : new Float32Array(tensor.data);
    const dims = tensor.dims;

    // Typical dims: [1, seq_len, hidden_size] or [seq_len, hidden_size]
    let seqLen: number;
    let hiddenSize: number;

    if (dims.length === 3) {
      // [batch=1, seq_len, hidden_size]
      seqLen = dims[1];
      hiddenSize = dims[2];
    } else if (dims.length === 2) {
      // [seq_len, hidden_size]
      seqLen = dims[0];
      hiddenSize = dims[1];
    } else {
      // Already flat — return as-is
      return data instanceof Float32Array ? data : new Float32Array(data);
    }

    // Mean pool over sequence dimension
    const result = new Float32Array(hiddenSize);
    for (let s = 0; s < seqLen; s++) {
      for (let h = 0; h < hiddenSize; h++) {
        result[h] += data[s * hiddenSize + h];
      }
    }
    for (let h = 0; h < hiddenSize; h++) {
      result[h] /= seqLen;
    }
    return result;
  }

  // Fallback: try tolist()
  if (tensor && typeof tensor.tolist === "function") {
    const list = tensor.tolist();
    if (Array.isArray(list) && Array.isArray(list[0])) {
      // [batch=1][seq_len][hidden]
      const inner = list[0];
      const hiddenSize2 = (inner[0] as number[]).length;
      const result2 = new Float32Array(hiddenSize2);
      for (const row of inner) {
        for (let h = 0; h < hiddenSize2; h++) {
          result2[h] += (row as number[])[h];
        }
      }
      for (let h = 0; h < hiddenSize2; h++) {
        result2[h] /= inner.length;
      }
      return result2;
    }
  }

  throw new Error("Cannot convert transformers output to Float32Array");
}

export async function embedTransformers(text: string): Promise<EmbeddingResult | EmbeddingError> {
  const t0 = Date.now();
  let pipe: PipelineFn;

  try {
    pipe = await getOrLoadPipeline();
  } catch (err: unknown) {
    const isTimeout =
      err instanceof Error &&
      (err.name === "AbortError" || err.message.toLowerCase().includes("timeout"));
    return {
      source: "transformers",
      model: TRANSFORMERS_MODEL,
      reason: isTimeout ? "timeout" : "model_load_failed",
      message: sanitizeErrorMessage(err instanceof Error ? err.message : String(err)),
    };
  }

  try {
    const output = await pipe(text, { pooling: "mean", normalize: true });
    const vector = tensorToFloat32Array(output);
    return {
      vector,
      source: "transformers",
      model: TRANSFORMERS_MODEL,
      dimensions: vector.length,
      latencyMs: Date.now() - t0,
      cached: false,
    };
  } catch (err: unknown) {
    const isTimeout =
      err instanceof Error &&
      (err.name === "AbortError" || err.message.toLowerCase().includes("timeout"));
    return {
      source: "transformers",
      model: TRANSFORMERS_MODEL,
      reason: isTimeout ? "timeout" : "request_failed",
      message: sanitizeErrorMessage(err instanceof Error ? err.message : String(err)),
    };
  }
}