@huggingface/inference
Version:
Typescript client for the Hugging Face Inference Providers and Inference Endpoints
440 lines (439 loc) • 18.6 kB
JavaScript
import { delay } from "../utils/delay.js";
import { omit } from "../utils/omit.js";
import { dataUrlFromBlob } from "../utils/dataUrlFromBlob.js";
import { BaseConversationalTask, BaseTextGenerationTask, TaskProviderHelper, } from "./providerHelper.js";
import { InferenceClientInputError, InferenceClientProviderApiError, InferenceClientProviderOutputError, } from "../errors.js";
const TOGETHER_API_BASE_URL = "https://api.together.xyz";
const AUDIO_MIME_TO_EXT = {
"audio/wav": "wav",
"audio/x-wav": "wav",
"audio/wave": "wav",
"audio/mpeg": "mp3",
"audio/mp3": "mp3",
"audio/mp4": "mp4",
"audio/m4a": "m4a",
"audio/x-m4a": "m4a",
"audio/flac": "flac",
"audio/x-flac": "flac",
"audio/ogg": "ogg",
"audio/webm": "webm",
};
function mimeTypeToExtension(mimeType) {
if (!mimeType) {
return "wav";
}
return AUDIO_MIME_TO_EXT[mimeType.toLowerCase()] ?? "wav";
}
export class TogetherConversationalTask extends BaseConversationalTask {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
preparePayload(params) {
const payload = super.preparePayload(params);
const response_format = payload.response_format;
if (response_format?.type === "json_schema" && response_format?.json_schema?.schema) {
payload.response_format = {
type: "json_schema",
schema: response_format.json_schema.schema,
};
}
return payload;
}
}
export class TogetherTextGenerationTask extends BaseTextGenerationTask {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
preparePayload(params) {
return {
model: params.model,
...params.args,
prompt: params.args.inputs,
};
}
async getResponse(response) {
if (typeof response === "object" &&
"choices" in response &&
Array.isArray(response?.choices) &&
typeof response?.model === "string") {
const completion = response.choices[0];
return {
generated_text: completion.text,
details: {
finish_reason: completion.finish_reason,
seed: completion.seed,
},
};
}
throw new InferenceClientProviderOutputError("Received malformed response from Together text generation API");
}
}
export class TogetherTextToImageTask extends TaskProviderHelper {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
makeRoute() {
return "v1/images/generations";
}
preparePayload(params) {
const rawParameters = params.args.parameters ?? {};
const { num_inference_steps, ...restParameters } = rawParameters;
if (num_inference_steps !== undefined) {
restParameters.steps = num_inference_steps;
}
return {
...omit(params.args, ["inputs", "parameters"]),
...restParameters,
prompt: params.args.inputs,
response_format: params.outputType === "url" ? "url" : "base64",
model: params.model,
};
}
/** Task label used in malformed-response errors. Overridden by subclasses. */
get imageTaskLabel() {
return "text-to-image";
}
async getResponse(response, url, headers, outputType, signal) {
void url;
void headers;
if (typeof response === "object" &&
"data" in response &&
Array.isArray(response.data) &&
response.data.length > 0) {
if (outputType === "json") {
return { ...response };
}
if ("url" in response.data[0] && typeof response.data[0].url === "string") {
return response.data[0].url;
}
if ("b64_json" in response.data[0] && typeof response.data[0].b64_json === "string") {
const base64Data = response.data[0].b64_json;
if (outputType === "dataUrl") {
return `data:image/jpeg;base64,${base64Data}`;
}
return fetch(`data:image/jpeg;base64,${base64Data}`, { signal }).then((res) => res.blob());
}
}
throw new InferenceClientProviderOutputError(`Received malformed response from Together ${this.imageTaskLabel} API`);
}
}
export class TogetherImageToImageTask extends TogetherTextToImageTask {
get imageTaskLabel() {
return "image-to-image";
}
preparePayload(params) {
const rawParameters = params.args.parameters ?? {};
const { prompt, num_inference_steps, ...restParameters } = rawParameters;
if (num_inference_steps !== undefined) {
restParameters.steps = num_inference_steps;
}
// Together exposes two mutually-exclusive image inputs (see
// https://docs.together.ai/docs/image-to-image): FLUX.1 Kontext only accepts
// `image_url`; FLUX.2 [dev] and Google models (Gemini 3 Pro Image, Flash Image
// 2.5) only accept `reference_images`. FLUX.2 [pro]/[flex] accept either but
// `reference_images` is the documented default. Use `image_url` only for
// FLUX.1 Kontext models and `reference_images` for everything else.
const lowered = params.model.toLowerCase();
const useImageUrl = lowered.includes("kontext") && lowered.includes("flux.1");
const imageField = useImageUrl
? { image_url: params.args.inputs }
: { reference_images: [params.args.inputs] };
return {
...omit(params.args, ["inputs", "parameters"]),
prompt: prompt ?? "",
...imageField,
...restParameters,
response_format: "base64",
model: params.model,
};
}
async preparePayloadAsync(args) {
const { inputs, ...restArgs } = args;
if (!(inputs instanceof Blob)) {
throw new InferenceClientInputError("Together image-to-image expects a Blob input.");
}
const imageDataUrl = await dataUrlFromBlob(inputs, inputs.type || "image/jpeg");
return {
...restArgs,
inputs: imageDataUrl,
};
}
async getResponse(response, url, headers, outputType) {
const result = await super.getResponse(response, url, headers, outputType);
if (result instanceof Blob) {
return result;
}
throw new InferenceClientProviderOutputError(`Received malformed response from Together ${this.imageTaskLabel} API`);
}
}
// Polling cadence for Together's async video generation.
const TOGETHER_VIDEO_POLLING_INTERVAL_MS = 2000;
// Upper bound on status polls (~5 minutes at TOGETHER_VIDEO_POLLING_INTERVAL_MS).
const TOGETHER_VIDEO_MAX_POLL_ATTEMPTS = 150;
// Statuses that mean "keep polling". Together returns "queued" before transitioning
// to "in_progress"; anything outside this set is treated as terminal.
const TOGETHER_VIDEO_PENDING_STATUSES = new Set(["queued", "in_progress"]);
/** Renames HF-standard fields to Together's video API field names. */
function normalizeTogetherVideoParameters(parameters) {
const { num_inference_steps, target_size, ...rest } = (parameters ?? {});
if (num_inference_steps !== undefined) {
rest.steps = num_inference_steps;
}
if (target_size && typeof target_size === "object") {
if (target_size.width !== undefined) {
rest.width = target_size.width;
}
if (target_size.height !== undefined) {
rest.height = target_size.height;
}
}
return rest;
}
/** Shared base for Together's async video tasks (text-to-video, image-to-video). */
class TogetherVideoTask extends TaskProviderHelper {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
makeRoute() {
return "v2/videos";
}
async getResponse(response, url, headers, _outputType, signal) {
if (!url || !headers) {
throw new InferenceClientInputError("URL and headers are required for Together video tasks");
}
const jobId = response?.id;
if (!jobId) {
throw new InferenceClientProviderOutputError("Received malformed response from Together video API: no job ID found in the response");
}
const statusUrl = `${url}/${jobId}`;
let job = response;
let status = job.status;
let attempt = 0;
// Together usually returns status: "queued" on the initial POST, but the field is
// typed as optional — treat a missing status as "pending" and poll, rather than
// falling through to the "unexpected status" error.
while (status === undefined || TOGETHER_VIDEO_PENDING_STATUSES.has(status)) {
if (attempt >= TOGETHER_VIDEO_MAX_POLL_ATTEMPTS) {
throw new InferenceClientProviderOutputError(`Timed out while waiting for Together video generation — aborting after ${TOGETHER_VIDEO_MAX_POLL_ATTEMPTS} status polls`);
}
attempt += 1;
await delay(TOGETHER_VIDEO_POLLING_INTERVAL_MS, signal);
const pollResponse = await fetch(statusUrl, { headers, signal });
if (!pollResponse.ok) {
throw new InferenceClientProviderApiError("Failed to fetch Together video job result", { url: statusUrl, method: "GET", headers }, {
requestId: pollResponse.headers.get("x-request-id") ?? "",
status: pollResponse.status,
body: await pollResponse.text(),
});
}
try {
job = (await pollResponse.json());
}
catch {
throw new InferenceClientProviderOutputError("Received malformed response from Together video API: failed to parse job result");
}
status = job.status;
}
if (status === "failed") {
throw new InferenceClientProviderOutputError(`Together video generation failed: ${job.error?.message ?? "Unknown error"}`);
}
if (status !== "completed") {
throw new InferenceClientProviderOutputError(`Unexpected Together video job status: ${JSON.stringify(status)}`);
}
const videoUrl = job.outputs?.video_url;
if (typeof videoUrl !== "string") {
throw new InferenceClientProviderOutputError("No video URL found in completed Together video job.");
}
const videoResponse = await fetch(videoUrl, { signal });
if (!videoResponse.ok) {
throw new InferenceClientProviderApiError("Failed to download Together video output", { url: videoUrl, method: "GET" }, {
requestId: videoResponse.headers.get("x-request-id") ?? "",
status: videoResponse.status,
body: await videoResponse.text(),
});
}
return await videoResponse.blob();
}
}
export class TogetherTextToVideoTask extends TogetherVideoTask {
preparePayload(params) {
return {
...omit(params.args, ["inputs", "parameters"]),
...normalizeTogetherVideoParameters(params.args.parameters),
prompt: params.args.inputs,
model: params.model,
};
}
}
export class TogetherImageToVideoTask extends TogetherVideoTask {
preparePayload(params) {
const rawParameters = params.args.parameters ?? {};
const { prompt, ...rest } = rawParameters;
const normalized = normalizeTogetherVideoParameters(rest);
const payload = {
...omit(params.args, ["inputs", "parameters"]),
...normalized,
// Together expects each keyframe as { input_image: <base64>, frame: "first" | "last" }
// for i2v models. See https://docs.together.ai/docs/inference/videos/reference-and-keyframes
frame_images: [{ input_image: params.args.inputs, frame: "first" }],
model: params.model,
};
if (typeof prompt === "string") {
payload.prompt = prompt;
}
return payload;
}
async preparePayloadAsync(args) {
const { inputs, ...restArgs } = args;
if (!(inputs instanceof Blob)) {
throw new InferenceClientInputError("Together image-to-video expects a Blob input.");
}
// Together's i2v models accept the image as a data URL or as an HTTP(S) URL in
// `frame_images[].input_image`, but cap the field at ~60KB. Larger images will
// be rejected with "Range of input length should be [1, 61440]" — users with
// big inputs should host the image and pass parameters.frame_images directly.
const imageDataUrl = await dataUrlFromBlob(inputs, inputs.type || "image/png");
return {
...restArgs,
inputs: imageDataUrl,
};
}
}
export class TogetherFeatureExtractionTask extends TaskProviderHelper {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
makeRoute() {
return "v1/embeddings";
}
preparePayload(params) {
return {
...omit(params.args, ["inputs", "parameters"]),
...params.args.parameters,
input: params.args.inputs,
model: params.model,
};
}
async getResponse(response) {
if (typeof response === "object" &&
response !== null &&
"data" in response &&
Array.isArray(response.data) &&
response.data.every((item) => typeof item === "object" && !!item && Array.isArray(item.embedding))) {
return response.data.map((item) => item.embedding);
}
throw new InferenceClientProviderOutputError(`Received malformed response from Together feature-extraction (embeddings) API: ${JSON.stringify(response)}`);
}
}
export class TogetherTextToSpeechTask extends TaskProviderHelper {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
makeRoute() {
return "v1/audio/speech";
}
preparePayload(params) {
const userParams = params.args.parameters ?? {};
// Together's /v1/audio/speech requires a `voice` field. Voices are model-specific
// (Kokoro accepts `af_*`, Orpheus uses different names, etc.), so we only default
// when the target model is Kokoro — the only TTS model currently registered.
const isKokoro = params.model.toLowerCase().includes("kokoro");
const voice = userParams.voice ?? (isKokoro ? "af_alloy" : undefined);
return {
...omit(params.args, ["inputs", "parameters"]),
...userParams,
...(voice !== undefined ? { voice } : {}),
input: params.args.inputs,
model: params.model,
};
}
async getResponse(response) {
if (response instanceof Blob) {
return response;
}
throw new InferenceClientProviderOutputError(`Received malformed response from Together text-to-speech API: ${JSON.stringify(response)}`);
}
}
export class TogetherAutomaticSpeechRecognitionTask extends TaskProviderHelper {
constructor() {
super("together", TOGETHER_API_BASE_URL);
}
makeRoute() {
return "v1/audio/transcriptions";
}
preparePayload(params) {
return {
...omit(params.args, ["inputs", "parameters", "data"]),
...params.args.parameters,
model: params.model,
};
}
makeBody(params) {
const audio = params.args.data;
const formData = new FormData();
if (audio instanceof Blob) {
formData.append("file", audio, `audio.${mimeTypeToExtension(audio.type)}`);
}
else if (typeof audio === "string") {
// Together's transcriptions endpoint also accepts a public HTTP(S) URL as the
// `file` form field (see https://docs.together.ai/docs/speech-to-text).
formData.append("file", audio);
}
else {
throw new InferenceClientInputError("Together automatic-speech-recognition expects a Blob, ArrayBuffer, or HTTP(S) URL string audio input.");
}
const fields = this.preparePayload(params);
for (const [key, value] of Object.entries(fields)) {
if (value === undefined || value === null) {
continue;
}
if (typeof value === "string") {
formData.append(key, value);
}
else if (typeof value === "number" || typeof value === "boolean") {
formData.append(key, String(value));
}
else {
formData.append(key, JSON.stringify(value));
}
}
return formData;
}
async preparePayloadAsync(args) {
const audio = "data" in args ? args.data : args.inputs;
let data;
if (audio instanceof Blob) {
data = audio;
}
else if (audio instanceof ArrayBuffer) {
data = new Blob([audio]);
}
else if (typeof audio === "string" && /^https?:\/\//.test(audio)) {
// Pass HTTP(S) URLs through as a string; makeBody will append them to the
// `file` form field instead of uploading bytes.
data = audio;
}
else {
throw new InferenceClientInputError("Together automatic-speech-recognition expects a Blob, ArrayBuffer, or HTTP(S) URL string audio input.");
}
// `data` is typed as `Blob | ArrayBuffer` in RequestArgs; URL-string is a
// Together-specific extension carried through to makeBody.
return {
...("data" in args ? omit(args, "data") : omit(args, "inputs")),
data,
};
}
async getResponse(response) {
if (typeof response === "object" && response !== null && typeof response.text === "string") {
const out = { text: response.text };
if (Array.isArray(response.segments)) {
out.chunks = response.segments.map((seg) => ({
text: seg.text,
timestamp: [seg.start, seg.end],
}));
}
return out;
}
throw new InferenceClientProviderOutputError(`Received malformed response from Together automatic-speech-recognition API: ${JSON.stringify(response)}`);
}
}