Team Ai
Apppublic

LiquidAI/LFM2-WebGPU

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
94likes
useLLM.ts271 linesDownload Raw Back to hooks
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};