feat: allow custom seed (#6709)

This commit is contained in:
darkskygit
2024-04-26 11:40:07 +00:00
parent 5d114ea965
commit b639e52dca
4 changed files with 59 additions and 79 deletions

View File

@@ -100,6 +100,17 @@ export class CopilotController {
return controller.signal; return controller.signal;
} }
private parseNumber(value: string | string[] | undefined) {
if (!value) {
return undefined;
}
const num = Number.parseInt(Array.isArray(value) ? value[0] : value, 10);
if (Number.isNaN(num)) {
return undefined;
}
return num;
}
private handleError(err: any) { private handleError(err: any) {
if (err instanceof Error) { if (err instanceof Error) {
const ret = { const ret = {
@@ -256,6 +267,7 @@ export class CopilotController {
return from( return from(
provider.generateImagesStream(session.finish(params), session.model, { provider.generateImagesStream(session.finish(params), session.model, {
seed: this.parseNumber(params.seed),
signal: this.getSignal(req), signal: this.getSignal(req),
user: user.id, user: user.id,
}) })

View File

@@ -2,6 +2,7 @@ import assert from 'node:assert';
import { import {
CopilotCapability, CopilotCapability,
CopilotImageOptions,
CopilotImageToImageProvider, CopilotImageToImageProvider,
CopilotProviderType, CopilotProviderType,
CopilotTextToImageProvider, CopilotTextToImageProvider,
@@ -57,10 +58,7 @@ export class FalProvider
async generateImages( async generateImages(
messages: PromptMessage[], messages: PromptMessage[],
model: string = this.availableModels[0], model: string = this.availableModels[0],
options: { options: CopilotImageOptions = {}
signal?: AbortSignal;
user?: string;
} = {}
): Promise<Array<string>> { ): Promise<Array<string>> {
const { content, attachments } = messages.pop() || {}; const { content, attachments } = messages.pop() || {};
if (!this.availableModels.includes(model)) { if (!this.availableModels.includes(model)) {
@@ -82,7 +80,7 @@ export class FalProvider
image_url: attachments?.[0], image_url: attachments?.[0],
prompt: content, prompt: content,
sync_mode: true, sync_mode: true,
seed: 42, seed: options.seed || 42,
enable_safety_checks: false, enable_safety_checks: false,
}), }),
signal: options.signal, signal: options.signal,
@@ -100,10 +98,7 @@ export class FalProvider
async *generateImagesStream( async *generateImagesStream(
messages: PromptMessage[], messages: PromptMessage[],
model: string = this.availableModels[0], model: string = this.availableModels[0],
options: { options: CopilotImageOptions = {}
signal?: AbortSignal;
user?: string;
} = {}
): AsyncIterable<string> { ): AsyncIterable<string> {
const ret = await this.generateImages(messages, model, options); const ret = await this.generateImages(messages, model, options);
for (const url of ret) { for (const url of ret) {

View File

@@ -5,6 +5,9 @@ import { ClientOptions, OpenAI } from 'openai';
import { import {
ChatMessageRole, ChatMessageRole,
CopilotCapability, CopilotCapability,
CopilotChatOptions,
CopilotEmbeddingOptions,
CopilotImageOptions,
CopilotImageToTextProvider, CopilotImageToTextProvider,
CopilotProviderType, CopilotProviderType,
CopilotTextToEmbeddingProvider, CopilotTextToEmbeddingProvider,
@@ -147,12 +150,7 @@ export class OpenAIProvider
async generateText( async generateText(
messages: PromptMessage[], messages: PromptMessage[],
model: string = 'gpt-3.5-turbo', model: string = 'gpt-3.5-turbo',
options: { options: CopilotChatOptions = {}
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
} = {}
): Promise<string> { ): Promise<string> {
this.checkParams({ messages, model }); this.checkParams({ messages, model });
const result = await this.instance.chat.completions.create( const result = await this.instance.chat.completions.create(
@@ -175,12 +173,7 @@ export class OpenAIProvider
async *generateTextStream( async *generateTextStream(
messages: PromptMessage[], messages: PromptMessage[],
model: string = 'gpt-3.5-turbo', model: string = 'gpt-3.5-turbo',
options: { options: CopilotChatOptions = {}
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
} = {}
): AsyncIterable<string> { ): AsyncIterable<string> {
this.checkParams({ messages, model }); this.checkParams({ messages, model });
const result = await this.instance.chat.completions.create( const result = await this.instance.chat.completions.create(
@@ -214,11 +207,7 @@ export class OpenAIProvider
async generateEmbedding( async generateEmbedding(
messages: string | string[], messages: string | string[],
model: string, model: string,
options: { options: CopilotEmbeddingOptions = { dimensions: DEFAULT_DIMENSIONS }
dimensions: number;
signal?: AbortSignal;
user?: string;
} = { dimensions: DEFAULT_DIMENSIONS }
): Promise<number[][]> { ): Promise<number[][]> {
messages = Array.isArray(messages) ? messages : [messages]; messages = Array.isArray(messages) ? messages : [messages];
this.checkParams({ embeddings: messages, model }); this.checkParams({ embeddings: messages, model });
@@ -236,10 +225,7 @@ export class OpenAIProvider
async generateImages( async generateImages(
messages: PromptMessage[], messages: PromptMessage[],
model: string = 'dall-e-3', model: string = 'dall-e-3',
options: { options: CopilotImageOptions = {}
signal?: AbortSignal;
user?: string;
} = {}
): Promise<Array<string>> { ): Promise<Array<string>> {
const { content: prompt } = messages.pop() || {}; const { content: prompt } = messages.pop() || {};
if (!prompt) { if (!prompt) {
@@ -261,10 +247,7 @@ export class OpenAIProvider
async *generateImagesStream( async *generateImagesStream(
messages: PromptMessage[], messages: PromptMessage[],
model: string = 'dall-e-3', model: string = 'dall-e-3',
options: { options: CopilotImageOptions = {}
signal?: AbortSignal;
user?: string;
} = {}
): AsyncIterable<string> { ): AsyncIterable<string> {
const ret = await this.generateImages(messages, model, options); const ret = await this.generateImages(messages, model, options);
for (const url of ret) { for (const url of ret) {

View File

@@ -143,6 +143,32 @@ export enum CopilotCapability {
ImageToText = 'image-to-text', ImageToText = 'image-to-text',
} }
const CopilotProviderOptionsSchema = z.object({
signal: z.instanceof(AbortSignal).optional(),
user: z.string().optional(),
});
const CopilotChatOptionsSchema = CopilotProviderOptionsSchema.extend({
temperature: z.number().optional(),
maxTokens: z.number().optional(),
}).optional();
export type CopilotChatOptions = z.infer<typeof CopilotChatOptionsSchema>;
const CopilotEmbeddingOptionsSchema = CopilotProviderOptionsSchema.extend({
dimensions: z.number(),
}).optional();
export type CopilotEmbeddingOptions = z.infer<
typeof CopilotEmbeddingOptionsSchema
>;
const CopilotImageOptionsSchema = CopilotProviderOptionsSchema.extend({
seed: z.number().optional(),
}).optional();
export type CopilotImageOptions = z.infer<typeof CopilotImageOptionsSchema>;
export interface CopilotProvider { export interface CopilotProvider {
readonly type: CopilotProviderType; readonly type: CopilotProviderType;
getCapabilities(): CopilotCapability[]; getCapabilities(): CopilotCapability[];
@@ -153,22 +179,12 @@ export interface CopilotTextToTextProvider extends CopilotProvider {
generateText( generateText(
messages: PromptMessage[], messages: PromptMessage[],
model?: string, model?: string,
options?: { options?: CopilotChatOptions
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
}
): Promise<string>; ): Promise<string>;
generateTextStream( generateTextStream(
messages: PromptMessage[], messages: PromptMessage[],
model?: string, model?: string,
options?: { options?: CopilotChatOptions
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
}
): AsyncIterable<string>; ): AsyncIterable<string>;
} }
@@ -176,11 +192,7 @@ export interface CopilotTextToEmbeddingProvider extends CopilotProvider {
generateEmbedding( generateEmbedding(
messages: string[] | string, messages: string[] | string,
model: string, model: string,
options: { options?: CopilotEmbeddingOptions
dimensions: number;
signal?: AbortSignal;
user?: string;
}
): Promise<number[][]>; ): Promise<number[][]>;
} }
@@ -188,18 +200,12 @@ export interface CopilotTextToImageProvider extends CopilotProvider {
generateImages( generateImages(
messages: PromptMessage[], messages: PromptMessage[],
model: string, model: string,
options: { options?: CopilotImageOptions
signal?: AbortSignal;
user?: string;
}
): Promise<Array<string>>; ): Promise<Array<string>>;
generateImagesStream( generateImagesStream(
messages: PromptMessage[], messages: PromptMessage[],
model?: string, model?: string,
options?: { options?: CopilotImageOptions
signal?: AbortSignal;
user?: string;
}
): AsyncIterable<string>; ): AsyncIterable<string>;
} }
@@ -207,22 +213,12 @@ export interface CopilotImageToTextProvider extends CopilotProvider {
generateText( generateText(
messages: PromptMessage[], messages: PromptMessage[],
model: string, model: string,
options: { options?: CopilotChatOptions
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
}
): Promise<string>; ): Promise<string>;
generateTextStream( generateTextStream(
messages: PromptMessage[], messages: PromptMessage[],
model: string, model: string,
options: { options?: CopilotChatOptions
temperature?: number;
maxTokens?: number;
signal?: AbortSignal;
user?: string;
}
): AsyncIterable<string>; ): AsyncIterable<string>;
} }
@@ -230,18 +226,12 @@ export interface CopilotImageToImageProvider extends CopilotProvider {
generateImages( generateImages(
messages: PromptMessage[], messages: PromptMessage[],
model: string, model: string,
options: { options?: CopilotImageOptions
signal?: AbortSignal;
user?: string;
}
): Promise<Array<string>>; ): Promise<Array<string>>;
generateImagesStream( generateImagesStream(
messages: PromptMessage[], messages: PromptMessage[],
model?: string, model?: string,
options?: { options?: CopilotImageOptions
signal?: AbortSignal;
user?: string;
}
): AsyncIterable<string>; ): AsyncIterable<string>;
} }