CoolFace
Apppublic

SpongeBobFan2002/openaudio-s1-mini

sourceHugging Facecc-by-nc-sa-4.0updated 1y agoView on Hugging Face
0likes
schema.py191 linesDownload Raw Back to tools
1import os
2import queue
3from dataclasses import dataclass
4from typing import Annotated, Literal, Optional
5
6import torch
7from pydantic import AfterValidator, BaseModel, Field, confloat, conint, conlist
8from pydantic.functional_validators import SkipValidation
9
10from fish_speech.conversation import Message, TextPart, VQPart
11
12GLOBAL_NUM_SAMPLES = int(os.getenv("GLOBAL_NUM_SAMPLES", 1))
13
14
15class ServeVQPart(BaseModel):
16    type: Literal["vq"] = "vq"
17    codes: SkipValidation[list[list[int]]]
18
19
20class ServeTextPart(BaseModel):
21    type: Literal["text"] = "text"
22    text: str
23
24
25class ServeAudioPart(BaseModel):
26    type: Literal["audio"] = "audio"
27    audio: bytes
28
29
30@dataclass
31class ASRPackRequest:
32    audio: torch.Tensor
33    result_queue: queue.Queue
34    language: str
35
36
37class ServeASRRequest(BaseModel):
38    # The audio should be an uncompressed PCM float16 audio
39    audios: list[bytes]
40    sample_rate: int = 44100
41    language: Literal["zh", "en", "ja", "auto"] = "auto"
42
43
44class ServeASRTranscription(BaseModel):
45    text: str
46    duration: float
47    huge_gap: bool
48
49
50class ServeASRSegment(BaseModel):
51    text: str
52    start: float
53    end: float
54
55
56class ServeTimedASRResponse(BaseModel):
57    text: str
58    segments: list[ServeASRSegment]
59    duration: float
60
61
62class ServeASRResponse(BaseModel):
63    transcriptions: list[ServeASRTranscription]
64
65
66class ServeMessage(BaseModel):
67    role: Literal["system", "assistant", "user", "raw"]
68    parts: list[ServeVQPart | ServeTextPart]
69
70    def to_conversation_message(self):
71        new_message = Message(role=self.role, parts=[])
72        if self.role == "assistant":
73            new_message.modality = "voice"
74
75        for part in self.parts:
76            if isinstance(part, ServeTextPart):
77                new_message.parts.append(TextPart(text=part.text))
78            elif isinstance(part, ServeVQPart):
79                new_message.parts.append(
80                    VQPart(codes=torch.tensor(part.codes, dtype=torch.int))
81                )
82            else:
83                raise ValueError(f"Unsupported part type: {part}")
84
85        return new_message
86
87
88class ServeRequest(BaseModel):
89    messages: Annotated[list[ServeMessage], conlist(ServeMessage, min_length=1)]
90    max_new_tokens: int = 1024
91    top_p: float = 0.7
92    repetition_penalty: float = 1.2
93    temperature: float = 0.7
94    streaming: bool = False
95    num_samples: int = 1
96    early_stop_threshold: float = 1.0
97
98
99class ServeVQGANEncodeRequest(BaseModel):
100    # The audio here should be in wav, mp3, etc
101    audios: list[bytes]
102
103
104class ServeVQGANEncodeResponse(BaseModel):
105    tokens: SkipValidation[list[list[list[int]]]]
106
107
108class ServeVQGANDecodeRequest(BaseModel):
109    tokens: SkipValidation[list[list[list[int]]]]
110
111
112class ServeVQGANDecodeResponse(BaseModel):
113    # The audio here should be in PCM float16 format
114    audios: list[bytes]
115
116
117class ServeReferenceAudio(BaseModel):
118    audio: bytes
119    text: str
120
121
122class ServeForwardMessage(BaseModel):
123    role: str
124    content: str
125
126
127class ServeResponse(BaseModel):
128    messages: list[ServeMessage]
129    finish_reason: Literal["stop", "error"] | None = None
130    stats: dict[str, int | float | str] = {}
131
132
133class ServeStreamDelta(BaseModel):
134    role: Literal["system", "assistant", "user"] | None = None
135    part: ServeVQPart | ServeTextPart | None = None
136
137
138class ServeStreamResponse(BaseModel):
139    sample_id: int = 0
140    delta: ServeStreamDelta | None = None
141    finish_reason: Literal["stop", "error"] | None = None
142    stats: dict[str, int | float | str] | None = None
143
144
145class ServeReferenceAudio(BaseModel):
146    audio: bytes
147    text: str
148
149    def __repr__(self) -> str:
150        return f"ServeReferenceAudio(text={self.text!r}, audio_size={len(self.audio)})"
151
152
153class ServeChatRequestV1(BaseModel):
154    model: str = "llama3-8b"
155    messages: list[ServeForwardMessage] = []
156    audio: bytes | None = None
157    temperature: float = 1.0
158    top_p: float = 1.0
159    max_tokens: int = 256
160    voice: str = "jessica"
161    tts_audio_format: Literal["mp3", "pcm", "opus"] = "mp3"
162    tts_audio_bitrate: Literal[16, 24, 32, 48, 64, 96, 128, 192] = 128
163
164
165class ServeTTSRequest(BaseModel):
166    text: str
167    chunk_length: Annotated[int, conint(ge=100, le=300, strict=True)] = 200
168    # Audio format
169    format: Literal["wav", "pcm", "mp3"] = "wav"
170    mp3_bitrate: Literal[64, 128, 192] = 128
171    # References audios for in-context learning
172    references: list[ServeReferenceAudio] = []
173    # Reference id
174    # For example, if you want use https://fish.audio/m/7f92f8afb8ec43bf81429cc1c9199cb1/
175    # Just pass 7f92f8afb8ec43bf81429cc1c9199cb1
176    reference_id: str | None = None
177    seed: int | None = None
178    use_memory_cache: Literal["on-demand", "never"] = "never"
179    # Normalize text for en & zh, this increase stability for numbers
180    normalize: bool = True
181    mp3_bitrate: Optional[int] = 64
182    opus_bitrate: Optional[int] = -1000
183    # Balance mode will reduce latency to 300ms, but may decrease stability
184    latency: Literal["normal", "balanced"] = "normal"
185    # not usually used below
186    streaming: bool = False
187    max_new_tokens: int = 1024
188    top_p: Annotated[float, Field(ge=0.1, le=1.0, strict=True)] = 0.7
189    repetition_penalty: Annotated[float, Field(ge=0.9, le=2.0, strict=True)] = 1.2
190    temperature: Annotated[float, Field(ge=0.1, le=1.0, strict=True)] = 0.7
191