LiquidAI/LFM2-WebGPU
94
1import { useState, useEffect, useRef, useCallback } from "react";2import {3 AutoModelForCausalLM,4 AutoTokenizer,5 TextStreamer,6} from "@huggingface/transformers";7 8// Define the supported model IDs9export type SupportedModelId = "350M" | "700M" | "1.2B";10 11export interface ModelConfig {12 dtype: string;13 device: string;14 revision?: string;15}16 17export const MODEL_CONFIGS: Record<SupportedModelId, ModelConfig> = {18 "350M": { dtype: "q4f16", device: "webgpu", revision: "5bc4b3e8cfd21660c0b1b9faa447ffbd9926b829" },19 "700M": { dtype: "q4f16", device: "webgpu", revision: "bf72eeabfe73a798674db899830a0dca99f8eabc" },20 "1.2B": { dtype: "q4f16", device: "webgpu", revision: "7f871660813dc1f34f0d304c77506c5fbdb440a0" },21};22 23interface LLMState {24 isLoading: boolean;25 isReady: boolean;26 error: string | null;27 progress: number;28}29 30interface LLMInstance {31 model: any;32 tokenizer: any;33}34 35let moduleCache: {36 [modelId: string]: {37 instance: LLMInstance | null;38 loadingPromise: Promise<LLMInstance> | null;39 };40} = {};41 42export const useLLM = (modelId?: SupportedModelId | string) => {43 const [state, setState] = useState<LLMState>({44 isLoading: false,45 isReady: false,46 error: null,47 progress: 0,48 });49 50 const instanceRef = useRef<LLMInstance | null>(null);51 const loadingPromiseRef = useRef<Promise<LLMInstance> | null>(null);52 53 const abortControllerRef = useRef<AbortController | null>(null);54 const pastKeyValuesRef = useRef<any>(null);55 56 const loadModel = useCallback(async () => {57 if (!modelId) {58 throw new Error("Model ID is required");59 }60 61 const MODEL_ID = `onnx-community/LFM2-${modelId}-ONNX`;62 63 if (!moduleCache[modelId]) {64 moduleCache[modelId] = {65 instance: null,66 loadingPromise: null,67 };68 }69 70 const cache = moduleCache[modelId];71 72 const existingInstance = instanceRef.current || cache.instance;73 if (existingInstance) {74 instanceRef.current = existingInstance;75 cache.instance = existingInstance;76 setState((prev) => ({ ...prev, isReady: true, isLoading: false }));77 return existingInstance;78 }79 80 const existingPromise = loadingPromiseRef.current || cache.loadingPromise;81 if (existingPromise) {82 try {83 const instance = await existingPromise;84 instanceRef.current = instance;85 cache.instance = instance;86 setState((prev) => ({ ...prev, isReady: true, isLoading: false }));87 return instance;88 } catch (error) {89 setState((prev) => ({90 ...prev,91 isLoading: false,92 error:93 error instanceof Error ? error.message : "Failed to load model",94 }));95 throw error;96 }97 }98 99 setState((prev) => ({100 ...prev,101 isLoading: true,102 error: null,103 progress: 0,104 }));105 106 abortControllerRef.current = new AbortController();107 108 const loadingPromise = (async () => {109 try {110 const progressCallback = (progress: any) => {111 // Only update progress for weights112 if (113 progress.status === "progress" &&114 progress.file.endsWith(".onnx_data")115 ) {116 const percentage = Math.round(117 (progress.loaded / progress.total) * 100,118 );119 setState((prev) => ({ ...prev, progress: percentage }));120 }121 };122 123 // Fallback to defaults if an unknown modelId string is passed124 const config = MODEL_CONFIGS[modelId as SupportedModelId] || {125 dtype: "q4f16",126 device: "webgpu",127 };128 129 const tokenizerOptions: Record<string, any> = {130 progress_callback: progressCallback,131 };132 if (config.revision) {133 tokenizerOptions.revision = config.revision;134 }135 136 const tokenizer = await AutoTokenizer.from_pretrained(137 MODEL_ID,138 tokenizerOptions139 );140 141 const modelOptions: Record<string, any> = {142 dtype: config.dtype,143 device: config.device,144 progress_callback: progressCallback,145 };146 if (config.revision) {147 modelOptions.revision = config.revision;148 }149 150 const model = await AutoModelForCausalLM.from_pretrained(151 MODEL_ID,152 modelOptions153 );154 155 const instance = { model, tokenizer };156 instanceRef.current = instance;157 cache.instance = instance;158 loadingPromiseRef.current = null;159 cache.loadingPromise = null;160 161 setState((prev) => ({162 ...prev,163 isLoading: false,164 isReady: true,165 progress: 100,166 }));167 return instance;168 } catch (error) {169 loadingPromiseRef.current = null;170 cache.loadingPromise = null;171 setState((prev) => ({172 ...prev,173 isLoading: false,174 error:175 error instanceof Error ? error.message : "Failed to load model",176 }));177 throw error;178 }179 })();180 181 loadingPromiseRef.current = loadingPromise;182 cache.loadingPromise = loadingPromise;183 return loadingPromise;184 }, [modelId]);185 186 const generateResponse = useCallback(187 async (188 messages: Array<{ role: string; content: string }>,189 tools: Array<any>,190 onToken?: (token: string) => void,191 ): Promise<string> => {192 const instance = instanceRef.current;193 if (!instance) {194 throw new Error("Model not loaded. Call loadModel() first.");195 }196 197 const { model, tokenizer } = instance;198 199 // Apply chat template with tools200 const input = tokenizer.apply_chat_template(messages, {201 tools,202 add_generation_prompt: true,203 return_dict: true,204 });205 206 const streamer = onToken207 ? new TextStreamer(tokenizer, {208 skip_prompt: true,209 skip_special_tokens: false,210 callback_function: (token: string) => {211 onToken(token);212 },213 })214 : undefined;215 216 // Generate the response217 const { sequences, past_key_values } = await model.generate({218 ...input,219 past_key_values: pastKeyValuesRef.current,220 max_new_tokens: 512,221 do_sample: false,222 streamer,223 return_dict_in_generate: true,224 });225 pastKeyValuesRef.current = past_key_values;226 227 // Decode the generated text with special tokens preserved (except final <|im_end|>) for tool call detection228 const response = tokenizer229 .batch_decode(sequences.slice(null, [input.input_ids.dims[1], null]), {230 skip_special_tokens: false,231 })[0]232 .replace(/<\|im_end\|>$/, "");233 234 return response;235 },236 [],237 );238 239 const clearPastKeyValues = useCallback(() => {240 pastKeyValuesRef.current = null;241 }, []);242 243 const cleanup = useCallback(() => {244 if (abortControllerRef.current) {245 abortControllerRef.current.abort();246 }247 }, []);248 249 useEffect(() => {250 return cleanup;251 }, [cleanup]);252 253 useEffect(() => {254 if (modelId && moduleCache[modelId]) {255 const existingInstance =256 instanceRef.current || moduleCache[modelId].instance;257 if (existingInstance) {258 instanceRef.current = existingInstance;259 setState((prev) => ({ ...prev, isReady: true }));260 }261 }262 }, [modelId]);263 264 return {265 ...state,266 loadModel,267 generateResponse,268 clearPastKeyValues,269 cleanup,270 };271};