#!/usr/bin/env bun /** * Download Embedding Models * * Pre-downloads all registered embedding models to ~/.cache/huggingface-transformers/ * This ensures models are available offline and avoids download time during benchmarks. * * Usage: * bun run models:download # Download all models * bun run models:download --small # Download only small models * bun run models:download --list # List available models */ import { pipeline, env } from '@huggingface/transformers'; import { homedir } from 'node:os'; import { join } from 'node:path'; import { MODEL_REGISTRY, RERANKER_REGISTRY, listModels } from '../src/models/registry.js'; import type { ModelConfig } from '../src/models/registry.js'; // Set cache directory env.cacheDir = join(homedir(), '.cache', 'huggingface-transformers'); interface DownloadResult { model: string; success: boolean; error?: string; timeMs: number; } async function downloadModel(config: ModelConfig): Promise { const startTime = Date.now(); try { console.log(` ↓ Downloading ${config.name} (${config.id})...`); // Load the model to trigger download const pipe = await pipeline('feature-extraction', config.id, { dtype: 'fp32', }); // Dispose to free memory await pipe.dispose(); const timeMs = Date.now() - startTime; console.log(` ✓ ${config.name} downloaded (${(timeMs / 1000).toFixed(1)}s)`); return { model: config.id, success: true, timeMs }; } catch (error) { const timeMs = Date.now() - startTime; const message = error instanceof Error ? error.message : String(error); console.log(` ✗ ${config.name} failed: ${message}`); return { model: config.id, success: false, error: message, timeMs }; } } async function downloadReranker(id: string, name: string): Promise { const startTime = Date.now(); try { console.log(` ↓ Downloading reranker ${name} (${id})...`); // Rerankers use text-classification pipeline const { AutoModelForSequenceClassification, AutoTokenizer } = await import( '@huggingface/transformers' ); const [model, tokenizer] = await Promise.all([ AutoModelForSequenceClassification.from_pretrained(id, { dtype: 'fp32' }), AutoTokenizer.from_pretrained(id), ]); // Dispose to free memory await model.dispose(); // Tokenizer doesn't have dispose const timeMs = Date.now() - startTime; console.log(` ✓ ${name} downloaded (${(timeMs / 1000).toFixed(1)}s)`); return { model: id, success: true, timeMs }; } catch (error) { const timeMs = Date.now() - startTime; const message = error instanceof Error ? error.message : String(error); console.log(` ✗ ${name} failed: ${message}`); return { model: id, success: false, error: message, timeMs }; } } function printModelList(): void { console.log('\nAvailable Embedding Models:\n'); const categories = ['bge', 'e5', 'minilm', 'gte', 'nomic', 'other'] as const; for (const category of categories) { const models = listModels({ category }); if (models.length === 0) continue; console.log(` ${category.toUpperCase()} Models:`); for (const model of models) { const size = model.sizeCategory.padEnd(5); const dims = String(model.dimensions).padStart(4); console.log(` [${size}] ${dims}d ${model.name}`); if (model.notes) { console.log(` ${model.notes}`); } } console.log(''); } console.log(' Reranker Models:'); for (const [key, config] of Object.entries(RERANKER_REGISTRY)) { console.log(` ${config.name} (${key})`); if (config.notes) { console.log(` ${config.notes}`); } } console.log(''); } async function main(): Promise { const args = process.argv.slice(2); if (args.includes('--list') || args.includes('-l')) { printModelList(); return; } const smallOnly = args.includes('--small') || args.includes('-s'); const skipRerankers = args.includes('--no-rerankers'); const specificModel = args.find((a) => !a.startsWith('-')); console.log('\n📦 Model Download Script\n'); console.log(`Cache directory: ${env.cacheDir}\n`); const results: DownloadResult[] = []; // Download embedding models let models: ModelConfig[]; if (specificModel !== undefined) { const config = MODEL_REGISTRY[specificModel]; if (config === undefined) { console.error(`Unknown model: ${specificModel}`); console.error(`Run with --list to see available models`); process.exit(1); } models = [config]; } else if (smallOnly) { models = listModels({ sizeCategory: 'small' }); console.log(`Downloading ${models.length} small embedding models...\n`); } else { models = Object.values(MODEL_REGISTRY); console.log(`Downloading ${models.length} embedding models...\n`); } for (const config of models) { const result = await downloadModel(config); results.push(result); } // Download rerankers if (!skipRerankers && specificModel === undefined) { console.log('\nDownloading reranker models...\n'); for (const [, config] of Object.entries(RERANKER_REGISTRY)) { const result = await downloadReranker(config.id, config.name); results.push(result); } } // Summary console.log('\n' + '─'.repeat(50)); const succeeded = results.filter((r) => r.success).length; const failed = results.filter((r) => !r.success).length; const totalTime = results.reduce((sum, r) => sum + r.timeMs, 0); console.log(`\n✓ ${succeeded} models downloaded`); if (failed > 0) { console.log(`✗ ${failed} models failed`); } console.log(`Total time: ${(totalTime / 1000).toFixed(1)}s\n`); if (failed > 0) { process.exit(1); } } main().catch((error) => { console.error('Fatal error:', error); process.exit(1); });