/** * HF-Inference do not have a mapping since all models use IDs from the Hub. * * If you want to try to run inference for a new model locally before it's registered on huggingface.co, * you can add it to the dictionary "HARDCODED_MODEL_ID_MAPPING" in consts.ts, for dev purposes. * * - If you work at HF and want to update this mapping, please use the model mapping API we provide on huggingface.co * - If you're a community member and want to add a new supported HF model to HF, please open an issue on the present repo * and we will tag HF team members. * * Thanks! */ import type { AudioClassificationOutput, AutomaticSpeechRecognitionOutput, ChatCompletionOutput, DocumentQuestionAnsweringOutput, FeatureExtractionOutput, FillMaskOutput, ImageClassificationOutput, ImageSegmentationOutput, ImageToTextOutput, ObjectDetectionOutput, QuestionAnsweringOutput, QuestionAnsweringOutputElement, SentenceSimilarityOutput, SummarizationOutput, TableQuestionAnsweringOutput, TextClassificationOutput, TextGenerationInput, TextGenerationOutput, TokenClassificationOutput, TranslationOutput, VisualQuestionAnsweringOutput, ZeroShotClassificationOutput, ZeroShotClassificationOutputElement, ZeroShotImageClassificationOutput, } from "@huggingface/tasks"; import { HF_ROUTER_URL } from "../config.js"; import { InferenceClientInputError, InferenceClientProviderOutputError } from "../errors.js"; import type { TabularClassificationOutput } from "../tasks/tabular/tabularClassification.js"; import type { BodyParams, OutputType, RequestArgs, UrlParams } from "../types.js"; import type { AudioClassificationTaskHelper, AudioToAudioTaskHelper, AutomaticSpeechRecognitionTaskHelper, ConversationalTaskHelper, DocumentQuestionAnsweringTaskHelper, FeatureExtractionTaskHelper, FillMaskTaskHelper, ImageClassificationTaskHelper, ImageSegmentationTaskHelper, ImageToImageTaskHelper, ImageToTextTaskHelper, ObjectDetectionTaskHelper, QuestionAnsweringTaskHelper, SentenceSimilarityTaskHelper, SummarizationTaskHelper, TableQuestionAnsweringTaskHelper, TabularClassificationTaskHelper, TabularRegressionTaskHelper, TextClassificationTaskHelper, TextGenerationTaskHelper, TextToAudioTaskHelper, TextToImageTaskHelper, TextToSpeechTaskHelper, TokenClassificationTaskHelper, TranslationTaskHelper, VisualQuestionAnsweringTaskHelper, ZeroShotClassificationTaskHelper, ZeroShotImageClassificationTaskHelper, } from "./providerHelper.js"; import { TaskProviderHelper } from "./providerHelper.js"; import { base64FromBytes } from "../utils/base64FromBytes.js"; import { dataUrlFromBlob } from "../utils/dataUrlFromBlob.js"; import type { ImageToImageArgs } from "../tasks/cv/imageToImage.js"; import type { AutomaticSpeechRecognitionArgs } from "../tasks/audio/automaticSpeechRecognition.js"; import type { AudioToAudioArgs } from "../tasks/audio/audioToAudio.js"; import { omit } from "../utils/omit.js"; import type { ImageSegmentationArgs } from "../tasks/cv/imageSegmentation.js"; import type { ImageToTextArgs } from "../tasks/cv/imageToText.js"; interface Base64ImageGeneration { data: Array<{ b64_json: string; }>; } interface OutputUrlImageGeneration { output: string[]; } interface AudioToAudioOutput { blob: string; "content-type": string; label: string; } export const EQUIVALENT_SENTENCE_TRANSFORMERS_TASKS = ["feature-extraction", "sentence-similarity"] as const; export class HFInferenceTask extends TaskProviderHelper { constructor() { super("hf-inference", `${HF_ROUTER_URL}/hf-inference`); } preparePayload(params: BodyParams): Record { return params.args; } override makeUrl(params: UrlParams): string { if (params.model.startsWith("http://") || params.model.startsWith("https://")) { return params.model; } return super.makeUrl(params); } makeRoute(params: UrlParams): string { if (params.task && ["feature-extraction", "sentence-similarity"].includes(params.task)) { // when deployed on hf-inference, those two tasks are automatically compatible with one another. return `models/${params.model}/pipeline/${params.task}`; } return `models/${params.model}`; } override async getResponse(response: unknown): Promise { return response; } } export class HFInferenceTextToImageTask extends HFInferenceTask implements TextToImageTaskHelper { override preparePayload(params: BodyParams): Record { if (params.outputType === "url") { throw new InferenceClientInputError( "hf-inference provider does not support URL output. Use outputType 'blob', 'dataUrl' or 'json' instead.", ); } return params.args; } override async getResponse( response: Base64ImageGeneration | OutputUrlImageGeneration, url?: string, headers?: HeadersInit, outputType?: OutputType, signal?: AbortSignal, ): Promise> { if (!response) { throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference text-to-image API: response is undefined", ); } if (typeof response == "object") { if (outputType === "json") { return { ...response }; } if ("data" in response && Array.isArray(response.data) && response.data[0].b64_json) { const base64Data = response.data[0].b64_json; if (outputType === "dataUrl") { return `data:image/jpeg;base64,${base64Data}`; } const base64Response = await fetch(`data:image/jpeg;base64,${base64Data}`, { signal }); return await base64Response.blob(); } if ("output" in response && Array.isArray(response.output)) { const urlResponse = await fetch(response.output[0], { signal }); const blob = await urlResponse.blob(); return outputType === "dataUrl" ? dataUrlFromBlob(blob) : blob; } } if (response instanceof Blob) { if (outputType === "dataUrl") { return dataUrlFromBlob(response); } if (outputType === "json") { return { output: await dataUrlFromBlob(response) }; } return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference text-to-image API: expected a Blob", ); } } export class HFInferenceConversationalTask extends HFInferenceTask implements ConversationalTaskHelper { override makeUrl(params: UrlParams): string { let url: string; if (params.model.startsWith("http://") || params.model.startsWith("https://")) { url = params.model.trim(); } else { url = `${this.makeBaseUrl(params)}/models/${params.model}`; } url = url.replace(/\/+$/, ""); if (url.endsWith("/v1")) { url += "/chat/completions"; } else if (!url.endsWith("/chat/completions")) { url += "/v1/chat/completions"; } return url; } override preparePayload(params: BodyParams): Record { return { ...params.args, model: params.model, }; } override async getResponse(response: ChatCompletionOutput): Promise { return response; } } interface HFInferenceTextCompletionOutput { choices: Array<{ text: string }>; } export class HFInferenceTextGenerationTask extends HFInferenceTask implements TextGenerationTaskHelper { override makeUrl(params: UrlParams): string { let url: string; if (params.model.startsWith("http://") || params.model.startsWith("https://")) { url = params.model.trim(); } else { url = `${this.makeBaseUrl(params)}/models/${params.model}`; } url = url.replace(/\/+$/, ""); if (url.endsWith("/v1")) { url += "/completions"; } else if (!url.endsWith("/completions")) { url += "/v1/completions"; } return url; } override preparePayload(params: BodyParams): Record { return { model: params.model, ...omit(params.args, ["inputs", "parameters"]), ...(params.args.parameters ? { max_tokens: params.args.parameters.max_new_tokens, ...omit(params.args.parameters, "max_new_tokens"), } : undefined), prompt: params.args.inputs, }; } override async getResponse(response: HFInferenceTextCompletionOutput): Promise { if ( typeof response === "object" && "choices" in response && Array.isArray(response.choices) && typeof response.choices[0]?.text === "string" ) { return { generated_text: response.choices[0].text }; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference text generation API: expected {choices: [{text: string}]}", ); } } export class HFInferenceAudioClassificationTask extends HFInferenceTask implements AudioClassificationTaskHelper { override async getResponse(response: unknown): Promise { if ( Array.isArray(response) && response.every( (x): x is { label: string; score: number } => typeof x === "object" && x !== null && typeof x.label === "string" && typeof x.score === "number", ) ) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference audio-classification API: expected Array<{label: string, score: number}> but received different format", ); } } export class HFInferenceAutomaticSpeechRecognitionTask extends HFInferenceTask implements AutomaticSpeechRecognitionTaskHelper { override async getResponse(response: AutomaticSpeechRecognitionOutput): Promise { return response; } async preparePayloadAsync(args: AutomaticSpeechRecognitionArgs): Promise { return "data" in args ? args : { ...omit(args, "inputs"), data: args.inputs, }; } } export class HFInferenceAudioToAudioTask extends HFInferenceTask implements AudioToAudioTaskHelper { async preparePayloadAsync(args: AudioToAudioArgs): Promise { return "data" in args ? args : { ...omit(args, "inputs"), data: args.inputs, }; } override async getResponse(response: AudioToAudioOutput[]): Promise { if (!Array.isArray(response)) { throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference audio-to-audio API: expected Array", ); } if ( !response.every((elem): elem is AudioToAudioOutput => { return ( typeof elem === "object" && elem && "label" in elem && typeof elem.label === "string" && "content-type" in elem && typeof elem["content-type"] === "string" && "blob" in elem && typeof elem.blob === "string" ); }) ) { throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference audio-to-audio API: expected Array<{label: string, audio: Blob}>", ); } return response; } } export class HFInferenceDocumentQuestionAnsweringTask extends HFInferenceTask implements DocumentQuestionAnsweringTaskHelper { override async getResponse( response: DocumentQuestionAnsweringOutput, ): Promise { if ( Array.isArray(response) && response.every( (elem) => typeof elem === "object" && !!elem && typeof elem?.answer === "string" && (typeof elem.end === "number" || typeof elem.end === "undefined") && (typeof elem.score === "number" || typeof elem.score === "undefined") && (typeof elem.start === "number" || typeof elem.start === "undefined"), ) ) { return response[0]; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference document-question-answering API: expected Array<{answer: string, end: number, score: number, start: number}>", ); } } export class HFInferenceFeatureExtractionTask extends HFInferenceTask implements FeatureExtractionTaskHelper { override async getResponse(response: FeatureExtractionOutput): Promise { const isNumArrayRec = (arr: unknown[], maxDepth: number, curDepth = 0): boolean => { if (curDepth > maxDepth) { return false; } if (arr.every((x) => Array.isArray(x))) { return arr.every((x) => isNumArrayRec(x as unknown[], maxDepth, curDepth + 1)); } else { return arr.every((x) => typeof x === "number"); } }; if (Array.isArray(response) && isNumArrayRec(response, 3, 0)) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference feature-extraction API: expected Array", ); } } export class HFInferenceImageClassificationTask extends HFInferenceTask implements ImageClassificationTaskHelper { override async getResponse(response: ImageClassificationOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x.label === "string" && typeof x.score === "number")) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference image-classification API: expected Array<{label: string, score: number}>", ); } } export class HFInferenceImageSegmentationTask extends HFInferenceTask implements ImageSegmentationTaskHelper { override async getResponse(response: ImageSegmentationOutput): Promise { if ( Array.isArray(response) && response.every( (x) => typeof x.label === "string" && typeof x.mask === "string" && (x.score === undefined || typeof x.score === "number"), ) ) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference image-segmentation API: expected Array<{label: string, mask: string, score: number}>", ); } async preparePayloadAsync(args: ImageSegmentationArgs): Promise { return { ...args, inputs: base64FromBytes( new Uint8Array(args.inputs instanceof ArrayBuffer ? args.inputs : await (args.inputs as Blob).arrayBuffer()), ), }; } } export class HFInferenceImageToTextTask extends HFInferenceTask implements ImageToTextTaskHelper { override async getResponse(response: ImageToTextOutput): Promise { if (typeof response?.generated_text !== "string") { throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference image-to-text API: expected {generated_text: string}", ); } return response; } async preparePayloadAsync(args: ImageToTextArgs): Promise { return "data" in args ? args : { ...omit(args, "inputs"), data: args.inputs }; } } export class HFInferenceImageToImageTask extends HFInferenceTask implements ImageToImageTaskHelper { async preparePayloadAsync(args: ImageToImageArgs): Promise { if (!args.parameters) { return { ...args, model: args.model, data: args.inputs, }; } else { return { ...args, inputs: base64FromBytes( new Uint8Array(args.inputs instanceof ArrayBuffer ? args.inputs : await (args.inputs as Blob).arrayBuffer()), ), }; } } override async getResponse(response: Blob): Promise { if (response instanceof Blob) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference image-to-image API: expected Blob", ); } } export class HFInferenceObjectDetectionTask extends HFInferenceTask implements ObjectDetectionTaskHelper { override async getResponse(response: ObjectDetectionOutput): Promise { if ( Array.isArray(response) && response.every( (x) => typeof x.label === "string" && typeof x.score === "number" && typeof x.box.xmin === "number" && typeof x.box.ymin === "number" && typeof x.box.xmax === "number" && typeof x.box.ymax === "number", ) ) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference object-detection API: expected Array<{label: string, score: number, box: {xmin: number, ymin: number, xmax: number, ymax: number}}>", ); } } export class HFInferenceZeroShotImageClassificationTask extends HFInferenceTask implements ZeroShotImageClassificationTaskHelper { override async getResponse(response: ZeroShotImageClassificationOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x.label === "string" && typeof x.score === "number")) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference zero-shot-image-classification API: expected Array<{label: string, score: number}>", ); } } export class HFInferenceTextClassificationTask extends HFInferenceTask implements TextClassificationTaskHelper { override async getResponse(response: TextClassificationOutput): Promise { const output = response?.[0]; if (Array.isArray(output) && output.every((x) => typeof x?.label === "string" && typeof x.score === "number")) { return output; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference text-classification API: expected Array<{label: string, score: number}>", ); } } export class HFInferenceQuestionAnsweringTask extends HFInferenceTask implements QuestionAnsweringTaskHelper { override async getResponse( response: QuestionAnsweringOutput | QuestionAnsweringOutput[number], ): Promise { if ( Array.isArray(response) ? response.every( (elem) => typeof elem === "object" && !!elem && typeof elem.answer === "string" && typeof elem.end === "number" && typeof elem.score === "number" && typeof elem.start === "number", ) : typeof response === "object" && !!response && typeof response.answer === "string" && typeof response.end === "number" && typeof response.score === "number" && typeof response.start === "number" ) { return Array.isArray(response) ? response[0] : response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference question-answering API: expected Array<{answer: string, end: number, score: number, start: number}>", ); } } export class HFInferenceFillMaskTask extends HFInferenceTask implements FillMaskTaskHelper { override async getResponse(response: FillMaskOutput): Promise { if ( Array.isArray(response) && response.every( (x) => typeof x.score === "number" && typeof x.sequence === "string" && typeof x.token === "number" && typeof x.token_str === "string", ) ) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference fill-mask API: expected Array<{score: number, sequence: string, token: number, token_str: string}>", ); } } export class HFInferenceZeroShotClassificationTask extends HFInferenceTask implements ZeroShotClassificationTaskHelper { override async getResponse(response: unknown): Promise { /// Handle Legacy response format from Inference API if ( typeof response === "object" && response !== null && "labels" in response && "scores" in response && Array.isArray(response.labels) && Array.isArray(response.scores) && response.labels.length === response.scores.length && response.labels.every((label: unknown): label is string => typeof label === "string") && response.scores.every((score: unknown): score is number => typeof score === "number") ) { const scores = response.scores; return response.labels.map((label: string, index: number) => ({ label, score: scores[index], })); } if (Array.isArray(response) && response.every(HFInferenceZeroShotClassificationTask.validateOutputElement)) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference zero-shot-classification API: expected Array<{label: string, score: number}>", ); } private static validateOutputElement(elem: unknown): elem is ZeroShotClassificationOutputElement { return ( typeof elem === "object" && !!elem && "label" in elem && "score" in elem && typeof elem.label === "string" && typeof elem.score === "number" ); } } export class HFInferenceSentenceSimilarityTask extends HFInferenceTask implements SentenceSimilarityTaskHelper { override async getResponse(response: SentenceSimilarityOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x === "number")) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference sentence-similarity API: expected Array", ); } } export class HFInferenceTableQuestionAnsweringTask extends HFInferenceTask implements TableQuestionAnsweringTaskHelper { static validate(elem: unknown): elem is TableQuestionAnsweringOutput[number] { return ( typeof elem === "object" && !!elem && "aggregator" in elem && typeof elem.aggregator === "string" && "answer" in elem && typeof elem.answer === "string" && "cells" in elem && Array.isArray(elem.cells) && elem.cells.every((x: unknown): x is string => typeof x === "string") && "coordinates" in elem && Array.isArray(elem.coordinates) && elem.coordinates.every( (coord: unknown): coord is number[] => Array.isArray(coord) && coord.every((x) => typeof x === "number"), ) ); } override async getResponse(response: TableQuestionAnsweringOutput): Promise { if ( Array.isArray(response) && Array.isArray(response) ? response.every((elem) => HFInferenceTableQuestionAnsweringTask.validate(elem)) : HFInferenceTableQuestionAnsweringTask.validate(response) ) { return Array.isArray(response) ? response[0] : response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference table-question-answering API: expected {aggregator: string, answer: string, cells: string[], coordinates: number[][]}", ); } } export class HFInferenceTokenClassificationTask extends HFInferenceTask implements TokenClassificationTaskHelper { override async getResponse(response: TokenClassificationOutput): Promise { if ( Array.isArray(response) && response.every( (x) => typeof x.end === "number" && typeof x.entity_group === "string" && typeof x.score === "number" && typeof x.start === "number" && typeof x.word === "string", ) ) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference token-classification API: expected Array<{end: number, entity_group: string, score: number, start: number, word: string}>", ); } } export class HFInferenceTranslationTask extends HFInferenceTask implements TranslationTaskHelper { override async getResponse(response: TranslationOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x?.translation_text === "string")) { return response?.length === 1 ? response?.[0] : response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference translation API: expected Array<{translation_text: string}>", ); } } export class HFInferenceSummarizationTask extends HFInferenceTask implements SummarizationTaskHelper { override async getResponse(response: SummarizationOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x?.summary_text === "string")) { return response?.[0]; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference summarization API: expected Array<{summary_text: string}>", ); } } export class HFInferenceTextToSpeechTask extends HFInferenceTask implements TextToSpeechTaskHelper { override async getResponse(response: Blob): Promise { return response; } } export class HFInferenceTabularClassificationTask extends HFInferenceTask implements TabularClassificationTaskHelper { override async getResponse(response: TabularClassificationOutput): Promise { if (Array.isArray(response) && response.every((x) => typeof x === "number")) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference tabular-classification API: expected Array", ); } } export class HFInferenceVisualQuestionAnsweringTask extends HFInferenceTask implements VisualQuestionAnsweringTaskHelper { override async getResponse(response: VisualQuestionAnsweringOutput): Promise { if ( Array.isArray(response) && response.every( (elem) => typeof elem === "object" && !!elem && typeof elem?.answer === "string" && typeof elem.score === "number", ) ) { return response[0]; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference visual-question-answering API: expected Array<{answer: string, score: number}>", ); } } export class HFInferenceTabularRegressionTask extends HFInferenceTask implements TabularRegressionTaskHelper { override async getResponse(response: number[]): Promise { if (Array.isArray(response) && response.every((x) => typeof x === "number")) { return response; } throw new InferenceClientProviderOutputError( "Received malformed response from HF-Inference tabular-regression API: expected Array", ); } } export class HFInferenceTextToAudioTask extends HFInferenceTask implements TextToAudioTaskHelper { override async getResponse(response: Blob): Promise { return response; } }