Spaces:
Running
Running
File size: 7,695 Bytes
95e5c44 | 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 | /**
* ToolTrace BiRefNet tracer backend (stateless image → mask).
* POST /trace body = crop image bytes (PNG/JPEG) → returns grayscale mask PNG
* GET /health → { ok, model, threads }
*
* BiRefNet captures chrome/reflective tools (the hammer shaft, caliper jaws) that
* IS-Net drops, but fp32/1024² OOMs the browser — so it runs here. Session is
* loaded once and warmed; threading + graph-opt minimise per-request latency.
*
* Run: node server/trace-server.cjs (local: uses root node_modules + local model)
* Env: PORT (8787) · BIREFNET_MODEL (path) · BIREFNET_MODEL_URL (fetch if missing) ·
* ORT_THREADS · TRACE_KEY (if set, /trace requires header X-Trace-Key to match)
*/
const http = require('http');
const os = require('os');
const fs = require('fs');
const path = require('path');
const https = require('https');
const httpc = require('http');
const ort = require('onnxruntime-node');
const sharp = require('sharp');
const PORT = Number(process.env.PORT) || 8787;
const MODEL = process.env.BIREFNET_MODEL || 'public/models/birefnet_lite.onnx';
const MODEL_URL = process.env.BIREFNET_MODEL_URL || ''; // download here if MODEL absent (Railway etc.)
const TRACE_KEY = process.env.TRACE_KEY || ''; // shared secret; empty = open (local dev)
const THREADS = Number(process.env.ORT_THREADS) || Math.max(1, os.cpus().length);
const SIZE = 1024;
const MEAN = [0.485, 0.456, 0.406], STD = [0.229, 0.224, 0.225];
let session = null;
// The 224MB model isn't in git. Locally it's a file on disk; on a host (Railway)
// set BIREFNET_MODEL_URL (e.g. a Cloudflare R2 / S3 object) and it's fetched once
// at startup. Follows redirects (object stores 302 to a signed URL).
function downloadModel(url, dest) {
return new Promise((resolve, reject) => {
fs.mkdirSync(path.dirname(dest), { recursive: true });
const tmp = `${dest}.part`;
const file = fs.createWriteStream(tmp);
const get = (u, depth) => {
if (depth > 5) { reject(new Error('too many redirects')); return; }
(u.startsWith('https') ? https : httpc).get(u, (res) => {
if ([301, 302, 303, 307, 308].includes(res.statusCode) && res.headers.location) { res.resume(); get(res.headers.location, depth + 1); return; }
if (res.statusCode !== 200) { res.resume(); reject(new Error(`model download HTTP ${res.statusCode}`)); return; }
res.pipe(file);
file.on('finish', () => file.close(() => { fs.renameSync(tmp, dest); resolve(); }));
}).on('error', reject);
};
get(url, 0);
});
}
async function init() {
const t0 = Date.now();
if (!fs.existsSync(MODEL)) {
if (!MODEL_URL) throw new Error(`Model not found at ${MODEL} and BIREFNET_MODEL_URL is not set`);
console.log(`model missing at ${MODEL} — downloading from ${MODEL_URL} …`);
await downloadModel(MODEL_URL, MODEL);
console.log(`model downloaded (${(fs.statSync(MODEL).size / 1e6).toFixed(0)}MB)`);
}
// Sequential + single inter-op thread, all intra-op threads on the one operator:
// fastest config in benchmarking (parallel/inter-op only added overhead here).
// Latency scales with physical cores — ~8s on a 4-core dev box, ~3-4s on a server.
session = await ort.InferenceSession.create(MODEL, {
intraOpNumThreads: THREADS,
interOpNumThreads: 1,
graphOptimizationLevel: 'all',
executionMode: 'sequential',
});
// Warm-up: first run JITs/allocates, so real requests don't pay that cost.
const warm = {}; warm[session.inputNames[0]] = new ort.Tensor('float32', new Float32Array(3 * SIZE * SIZE), [1, 3, SIZE, SIZE]);
const tw = Date.now(); await session.run(warm);
console.log(`BiRefNet ready: load ${((tw - t0) / 1000).toFixed(1)}s · warmup ${((Date.now() - tw) / 1000).toFixed(1)}s · threads=${THREADS}`);
}
async function trace(imgBuf) {
const meta = await sharp(imgBuf).metadata();
const W = meta.width, H = meta.height;
const { data } = await sharp(imgBuf).resize(SIZE, SIZE, { fit: 'fill' }).removeAlpha().raw().toBuffer({ resolveWithObject: true });
const plane = SIZE * SIZE;
const input = new Float32Array(3 * plane);
for (let i = 0; i < plane; i++) {
input[i] = (data[i * 3] / 255 - MEAN[0]) / STD[0];
input[plane + i] = (data[i * 3 + 1] / 255 - MEAN[1]) / STD[1];
input[2 * plane + i] = (data[i * 3 + 2] / 255 - MEAN[2]) / STD[2];
}
const feeds = {}; feeds[session.inputNames[0]] = new ort.Tensor('float32', input, [1, 3, SIZE, SIZE]);
const out = await session.run(feeds);
const sal = out[session.outputNames[0]].data;
let mn = Infinity, mx = -Infinity;
for (let i = 0; i < plane; i++) { const v = sal[i]; if (v < mn) mn = v; if (v > mx) mx = v; }
const range = mx - mn || 1;
const gray = Buffer.allocUnsafe(plane);
for (let i = 0; i < plane; i++) { const v = ((sal[i] - mn) / range) * 255; gray[i] = v < 0 ? 0 : v > 255 ? 255 : v; }
// resize the 1024² mask back to the crop's size; client thresholds + traces it
return sharp(gray, { raw: { width: SIZE, height: SIZE, channels: 1 } }).resize(W, H, { fit: 'fill' }).png().toBuffer();
}
const server = http.createServer((req, res) => {
res.setHeader('Access-Control-Allow-Origin', '*');
res.setHeader('Access-Control-Allow-Methods', 'POST, GET, OPTIONS');
res.setHeader('Access-Control-Allow-Headers', 'Content-Type, X-Trace-Key');
// The app is cross-origin isolated (COOP/COEP require-corp for WASM threads);
// this lets its fetch() consume our cross-origin response.
res.setHeader('Cross-Origin-Resource-Policy', 'cross-origin');
if (req.method === 'OPTIONS') { res.writeHead(204); res.end(); return; }
if (req.method === 'GET' && req.url === '/health') {
res.writeHead(200, { 'Content-Type': 'application/json' });
res.end(JSON.stringify({ ok: !!session, model: 'BiRefNet Lite', threads: THREADS }));
return;
}
// Shared-secret gate (only when TRACE_KEY is configured, so local dev stays open).
if (TRACE_KEY && req.headers['x-trace-key'] !== TRACE_KEY) { res.writeHead(401); res.end('unauthorized'); return; }
if (req.method === 'POST' && req.url === '/trace') {
if (!session) { res.writeHead(503); res.end('model not ready'); return; }
const chunks = [];
req.on('data', (c) => chunks.push(c));
req.on('end', async () => {
const t0 = Date.now();
try {
const maskPng = await trace(Buffer.concat(chunks));
res.writeHead(200, { 'Content-Type': 'image/png', 'X-Trace-Ms': String(Date.now() - t0) });
res.end(maskPng);
console.log(`/trace ${Date.now() - t0}ms (${Buffer.concat(chunks).length}b in)`);
} catch (e) { console.error('/trace error', e); res.writeHead(500); res.end(String(e.message || e)); }
});
return;
}
res.writeHead(404); res.end('not found');
});
// Bind the port FIRST (cheap), then load the model — so a port clash fails in
// milliseconds with a clear message instead of after the ~18s model warm-up.
// /health reports ok:false and /trace returns 503 until the session is ready.
server.on('error', (e) => {
if (e.code === 'EADDRINUSE') {
console.error(`Port ${PORT} is already in use — a tracer is probably already running.\n` +
`Stop the other instance first, or start this one on another port: PORT=8788 node server/trace-server.cjs`);
process.exit(1);
}
throw e;
});
server.listen(PORT, () => {
console.log(`ToolTrace tracer listening on http://localhost:${PORT} — loading model…`);
init().catch((e) => { console.error('init failed', e); process.exit(1); });
});
|