bupalinyu commited on
Commit
ceb8371
·
1 Parent(s): 888d54b
Files changed (6) hide show
  1. .gitignore +7 -0
  2. README.md +18 -3
  3. app.py +563 -0
  4. packages.txt +2 -0
  5. requirements.txt +9 -0
  6. scripts/run_local_gradio.sh +35 -0
.gitignore ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ __pycache__/
2
+ *.py[cod]
3
+ .gradio/
4
+ .venv/
5
+ venv/
6
+ runs/
7
+ tmp/
README.md CHANGED
@@ -4,10 +4,25 @@ emoji: 🐠
4
  colorFrom: green
5
  colorTo: indigo
6
  sdk: gradio
7
- sdk_version: 6.14.0
8
- python_version: '3.13'
9
  app_file: app.py
10
  pinned: false
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
  colorFrom: green
5
  colorTo: indigo
6
  sdk: gradio
7
+ sdk_version: 5.50.0
8
+ python_version: '3.10'
9
  app_file: app.py
10
  pinned: false
11
  ---
12
 
13
+ This Space runs an Ark-ASR 0.6B demo with the Transformers backend.
14
+
15
+ Default runtime limits are set for a single-GPU Space with GPU memory usage
16
+ targeted below 15 GB:
17
+
18
+ - `ARK_ASR_MAX_AUDIO_SECONDS=30`
19
+ - `ARK_ASR_DTYPE=float16`
20
+ - `ARK_ASR_ATTN_IMPL=sdpa`
21
+ - Gradio queue concurrency is limited to one request.
22
+
23
+ For local testing on this machine, use the local checkpoint instead of pulling
24
+ from the Hub:
25
+
26
+ ```bash
27
+ scripts/run_local_gradio.sh
28
+ ```
app.py ADDED
@@ -0,0 +1,563 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import logging
5
+ import os
6
+ import re
7
+ import tempfile
8
+ import time
9
+ from dataclasses import dataclass
10
+ from pathlib import Path
11
+ from typing import Any, Iterable
12
+
13
+ import gradio as gr
14
+ import torch
15
+ from transformers import AutoModelForCausalLM, AutoProcessor, AutoTokenizer
16
+ from transformers.generation.logits_process import LogitsProcessor, LogitsProcessorList
17
+
18
+
19
+ logging.basicConfig(
20
+ level=os.getenv("LOG_LEVEL", "INFO"),
21
+ format="[%(asctime)s] %(levelname)s %(name)s: %(message)s",
22
+ )
23
+ logger = logging.getLogger("ark_asr_space")
24
+
25
+ MODEL_ID = os.getenv("ARK_ASR_MODEL_ID", "AutoArk-AI/ARK-ASR-0.6B")
26
+ ASR_INSTRUCTION = os.getenv("ARK_ASR_INSTRUCTION", "Please transcribe this audio.")
27
+ MAX_AUDIO_SECONDS = int(os.getenv("ARK_ASR_MAX_AUDIO_SECONDS", "30"))
28
+ SAMPLING_RATE = int(os.getenv("ARK_ASR_SAMPLING_RATE", "16000"))
29
+ MAX_NEW_TOKENS = int(os.getenv("ARK_ASR_MAX_NEW_TOKENS", "256"))
30
+ DTYPE = os.getenv("ARK_ASR_DTYPE", "float16")
31
+ ATTN_IMPL = os.getenv("ARK_ASR_ATTN_IMPL", "sdpa")
32
+ ASR_BLOCK_TOKEN_ID_FROM = int(os.getenv("ARK_ASR_BLOCK_TOKEN_ID_FROM", "151670"))
33
+
34
+ SPECIAL_TOKEN_PATTERN = re.compile(
35
+ r"<\|(?:"
36
+ r"bicodec_(?:semantic|global)_\d+|"
37
+ r"(?:start|end)_(?:global_token|glm_token|semantic_token|content)"
38
+ r")\|>"
39
+ )
40
+ TURN_END_MARKERS = ("<|user|>", "<|assistant|>", "<|im_end|>")
41
+ LEADING_NOISE_PATTERN = re.compile(r"^[\s,.;:!?-]+")
42
+ CONTROL_TOKEN_PATTERN = re.compile(r"^<.*>$")
43
+
44
+
45
+ class BlockTokenIdsFromLogitsProcessor(LogitsProcessor):
46
+ def __init__(self, block_from_id: int | None, block_token_ids: Iterable[int] | None = None):
47
+ self.block_from_id = (
48
+ None if block_from_id is None or int(block_from_id) < 0 else int(block_from_id)
49
+ )
50
+ self.block_token_ids = sorted(set(int(token_id) for token_id in (block_token_ids or [])))
51
+
52
+ def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor:
53
+ vocab_size = scores.shape[-1]
54
+ if self.block_from_id is not None and self.block_from_id < vocab_size:
55
+ scores[:, self.block_from_id :] = -float("inf")
56
+ valid_token_ids = [token_id for token_id in self.block_token_ids if 0 <= token_id < vocab_size]
57
+ if valid_token_ids:
58
+ scores[:, valid_token_ids] = -float("inf")
59
+ return scores
60
+
61
+
62
+ @dataclass
63
+ class AppState:
64
+ model_path: str = ""
65
+ device: str = "cpu"
66
+ torch_dtype: torch.dtype = torch.float32
67
+ model: Any = None
68
+ processor: Any = None
69
+ tokenizer: Any = None
70
+ eos_token_ids: list[int] | None = None
71
+ extra_block_token_ids: list[int] | None = None
72
+ resolved_attn_impl: str = ""
73
+ loaded_at: float = 0.0
74
+
75
+
76
+ state = AppState()
77
+ load_lock = asyncio.Lock()
78
+ infer_lock = asyncio.Lock()
79
+
80
+
81
+ def normalize_token_ids(token_ids: Any) -> list[int]:
82
+ if token_ids is None:
83
+ return []
84
+ if isinstance(token_ids, (list, tuple, set)):
85
+ return [int(token_id) for token_id in token_ids if token_id is not None]
86
+ return [int(token_ids)]
87
+
88
+
89
+ def build_eos_token_ids(tokenizer: Any) -> list[int]:
90
+ eos_ids = []
91
+ eos_ids.extend(normalize_token_ids(getattr(tokenizer, "eos_token_id", None)))
92
+ for marker in TURN_END_MARKERS:
93
+ token_id = tokenizer.convert_tokens_to_ids(marker)
94
+ if isinstance(token_id, int) and token_id >= 0:
95
+ eos_ids.append(int(token_id))
96
+ return list(dict.fromkeys(eos_ids))
97
+
98
+
99
+ def build_asr_keep_token_ids(model: Any, tokenizer: Any) -> list[int]:
100
+ keep_token_ids = set()
101
+ keep_token_ids.update(normalize_token_ids(getattr(tokenizer, "eos_token_id", None)))
102
+ keep_token_ids.update(normalize_token_ids(getattr(getattr(model, "config", None), "eos_token_id", None)))
103
+ keep_token_ids.update(
104
+ normalize_token_ids(getattr(getattr(model, "generation_config", None), "eos_token_id", None))
105
+ )
106
+ return sorted(keep_token_ids)
107
+
108
+
109
+ def build_asr_extra_block_token_ids(
110
+ tokenizer: Any,
111
+ keep_token_ids: Iterable[int] | None = None,
112
+ block_from_id: int | None = None,
113
+ ) -> list[int]:
114
+ keep = set(int(token_id) for token_id in (keep_token_ids or []))
115
+ max_control_token_id = None if block_from_id is None or int(block_from_id) < 0 else int(block_from_id)
116
+ block_token_ids = {
117
+ int(token_id)
118
+ for token_id in getattr(tokenizer, "all_special_ids", [])
119
+ if token_id is not None
120
+ }
121
+ added_tokens_decoder = getattr(tokenizer, "added_tokens_decoder", {}) or {}
122
+ for token_id, token_meta in added_tokens_decoder.items():
123
+ token_id = int(token_id)
124
+ if max_control_token_id is not None and token_id >= max_control_token_id:
125
+ continue
126
+ token_content = getattr(token_meta, "content", None)
127
+ if token_content is None and isinstance(token_meta, dict):
128
+ token_content = token_meta.get("content")
129
+ if token_content and CONTROL_TOKEN_PATTERN.match(token_content):
130
+ block_token_ids.add(token_id)
131
+ block_token_ids.difference_update(keep)
132
+ return sorted(block_token_ids)
133
+
134
+
135
+ def truncate_generation_text(text: str) -> str:
136
+ if not text:
137
+ return ""
138
+ cut = len(text)
139
+ for marker in TURN_END_MARKERS:
140
+ index = text.find(marker)
141
+ if index != -1 and index < cut:
142
+ cut = index
143
+ return text[:cut].strip()
144
+
145
+
146
+ def remove_special_tokens(text: str) -> str:
147
+ if not text:
148
+ return ""
149
+ if "<|text|>" in text:
150
+ text = text.split("<|text|>", 1)[1]
151
+ return SPECIAL_TOKEN_PATTERN.sub("", text).strip()
152
+
153
+
154
+ def normalize_prediction_text(text: str) -> str:
155
+ if not text:
156
+ return ""
157
+ text = truncate_generation_text(text)
158
+ text = remove_special_tokens(text)
159
+ text = re.sub(r"\s+", " ", text).strip()
160
+ return LEADING_NOISE_PATTERN.sub("", text).strip()
161
+
162
+
163
+ def as_dict(value: Any) -> dict[str, Any]:
164
+ if isinstance(value, dict):
165
+ return value
166
+ if hasattr(value, "keys") and hasattr(value, "__getitem__"):
167
+ return {key: value[key] for key in value.keys()}
168
+ raise TypeError(f"Unexpected processor output type: {type(value)}")
169
+
170
+
171
+ def resolve_torch_dtype(dtype_name: str, device: str) -> torch.dtype:
172
+ if dtype_name == "auto":
173
+ return torch.float16 if device == "cuda" else torch.float32
174
+ mapping = {
175
+ "float16": torch.float16,
176
+ "bfloat16": torch.bfloat16,
177
+ "float32": torch.float32,
178
+ }
179
+ if dtype_name not in mapping:
180
+ raise ValueError(f"Unsupported dtype: {dtype_name}")
181
+ if device != "cuda" and mapping[dtype_name] != torch.float32:
182
+ return torch.float32
183
+ return mapping[dtype_name]
184
+
185
+
186
+ def maybe_gpu_memory_text() -> str:
187
+ if not torch.cuda.is_available():
188
+ return "GPU: not available in this runtime."
189
+ index = torch.cuda.current_device()
190
+ props = torch.cuda.get_device_properties(index)
191
+ total_gb = props.total_memory / 1024**3
192
+ reserved_gb = torch.cuda.memory_reserved(index) / 1024**3
193
+ allocated_gb = torch.cuda.memory_allocated(index) / 1024**3
194
+ return (
195
+ f"GPU: {props.name}, total={total_gb:.1f}G, "
196
+ f"reserved={reserved_gb:.1f}G, allocated={allocated_gb:.1f}G."
197
+ )
198
+
199
+
200
+ def resolve_model_path() -> str:
201
+ local_path = Path(MODEL_ID).expanduser()
202
+ if local_path.exists():
203
+ logger.info("Using local model path: %s", local_path.resolve())
204
+ return str(local_path.resolve())
205
+ logger.info("Using Hugging Face model id: %s", MODEL_ID)
206
+ return MODEL_ID
207
+
208
+
209
+ def load_model(model_path: str, device: str, torch_dtype: torch.dtype, attn_impl: str):
210
+ candidates = ["sdpa", "eager"] if attn_impl == "auto" else [attn_impl]
211
+ if attn_impl == "flash_attention_2":
212
+ candidates.extend(["sdpa", "eager"])
213
+
214
+ last_error: Exception | None = None
215
+ for candidate in candidates:
216
+ try:
217
+ logger.info("Loading model with attn_implementation=%s", candidate)
218
+ model = AutoModelForCausalLM.from_pretrained(
219
+ model_path,
220
+ trust_remote_code=True,
221
+ torch_dtype=torch_dtype,
222
+ attn_implementation=candidate,
223
+ ).to(device)
224
+ model.eval()
225
+ return model, candidate
226
+ except (ImportError, RuntimeError, ValueError) as exc:
227
+ if candidate != "flash_attention_2":
228
+ raise
229
+ logger.warning("flash_attention_2 unavailable, falling back: %s", str(exc).splitlines()[0])
230
+ last_error = exc
231
+ if last_error is not None:
232
+ raise last_error
233
+ raise RuntimeError("Failed to load model")
234
+
235
+
236
+ async def ensure_loaded() -> None:
237
+ if state.model is not None:
238
+ return
239
+
240
+ async with load_lock:
241
+ if state.model is not None:
242
+ return
243
+
244
+ started = time.perf_counter()
245
+ state.device = "cuda" if torch.cuda.is_available() else "cpu"
246
+ state.torch_dtype = resolve_torch_dtype(DTYPE, state.device)
247
+ state.model_path = resolve_model_path()
248
+
249
+ logger.info(
250
+ "Loading Transformers ASR stack: model=%s device=%s dtype=%s",
251
+ state.model_path,
252
+ state.device,
253
+ state.torch_dtype,
254
+ )
255
+ state.model, state.resolved_attn_impl = await asyncio.to_thread(
256
+ load_model,
257
+ state.model_path,
258
+ state.device,
259
+ state.torch_dtype,
260
+ ATTN_IMPL,
261
+ )
262
+ state.tokenizer = AutoTokenizer.from_pretrained(
263
+ state.model_path,
264
+ trust_remote_code=True,
265
+ fix_mistral_regex=True,
266
+ )
267
+ if state.tokenizer.pad_token_id is None:
268
+ state.tokenizer.pad_token_id = state.tokenizer.eos_token_id
269
+ state.tokenizer.padding_side = "left"
270
+
271
+ state.processor = AutoProcessor.from_pretrained(
272
+ state.model_path,
273
+ trust_remote_code=True,
274
+ fix_mistral_regex=True,
275
+ )
276
+ if hasattr(state.processor, "tokenizer"):
277
+ if state.processor.tokenizer.pad_token_id is None:
278
+ state.processor.tokenizer.pad_token_id = state.tokenizer.pad_token_id
279
+ state.processor.tokenizer.padding_side = "left"
280
+
281
+ state.eos_token_ids = build_eos_token_ids(state.tokenizer)
282
+ keep_token_ids = build_asr_keep_token_ids(state.model, state.tokenizer)
283
+ state.extra_block_token_ids = build_asr_extra_block_token_ids(
284
+ state.tokenizer,
285
+ keep_token_ids=keep_token_ids,
286
+ block_from_id=ASR_BLOCK_TOKEN_ID_FROM,
287
+ )
288
+ state.loaded_at = time.time()
289
+ logger.info(
290
+ "Transformers ASR stack loaded in %.2fs with attn=%s",
291
+ time.perf_counter() - started,
292
+ state.resolved_attn_impl,
293
+ )
294
+
295
+
296
+ def build_conversation(audio_path: str, begin_time: float, end_time: float) -> list[dict[str, Any]]:
297
+ return [
298
+ {
299
+ "role": "user",
300
+ "content": [
301
+ {
302
+ "type": "audio",
303
+ "path": audio_path,
304
+ "begin_time": begin_time,
305
+ "end_time": end_time,
306
+ },
307
+ {"type": "text", "text": ASR_INSTRUCTION},
308
+ ],
309
+ }
310
+ ]
311
+
312
+
313
+ def audio_to_path(audio: str | tuple[int, Any] | None) -> tuple[str, str | None]:
314
+ if audio is None:
315
+ raise gr.Error("Please upload or record an audio clip first.")
316
+ if isinstance(audio, str):
317
+ return audio, None
318
+
319
+ sample_rate, data = audio
320
+ tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".wav")
321
+ tmp.close()
322
+ import soundfile as sf
323
+
324
+ sf.write(tmp.name, data, sample_rate)
325
+ return tmp.name, tmp.name
326
+
327
+
328
+ def run_transformers_generation(
329
+ audio_path: str,
330
+ begin_time: float,
331
+ end_time: float,
332
+ max_new_tokens: int,
333
+ ) -> tuple[str, int]:
334
+ inputs_raw = state.processor.apply_chat_template(
335
+ [build_conversation(audio_path, begin_time, end_time)],
336
+ return_tensors="pt",
337
+ sampling_rate=SAMPLING_RATE,
338
+ audio_padding="longest",
339
+ add_generation_prompt=True,
340
+ text_kwargs={"padding": "longest"},
341
+ audio_max_length=int(MAX_AUDIO_SECONDS * SAMPLING_RATE),
342
+ )
343
+ if torch.is_tensor(inputs_raw):
344
+ raise RuntimeError("ASR apply_chat_template returned Tensor-only; audio was not encoded.")
345
+
346
+ inputs = as_dict(inputs_raw)
347
+ if "audios" not in inputs:
348
+ raise RuntimeError(f"ASR inputs missing 'audios'; processor keys={list(inputs.keys())}")
349
+ if "attention_mask" not in inputs and "input_ids" in inputs and torch.is_tensor(inputs["input_ids"]):
350
+ inputs["attention_mask"] = torch.ones_like(inputs["input_ids"], dtype=torch.long)
351
+
352
+ for key, value in list(inputs.items()):
353
+ if not torch.is_tensor(value):
354
+ continue
355
+ if key == "audios":
356
+ inputs[key] = value.to(device=state.device, dtype=state.torch_dtype)
357
+ else:
358
+ inputs[key] = value.to(state.device)
359
+
360
+ generate_kwargs: dict[str, Any] = {
361
+ "max_new_tokens": int(max_new_tokens),
362
+ "do_sample": False,
363
+ "pad_token_id": state.tokenizer.pad_token_id,
364
+ }
365
+ if state.eos_token_ids:
366
+ generate_kwargs["eos_token_id"] = state.eos_token_ids
367
+ if ASR_BLOCK_TOKEN_ID_FROM >= 0 or state.extra_block_token_ids:
368
+ generate_kwargs["logits_processor"] = LogitsProcessorList(
369
+ [
370
+ BlockTokenIdsFromLogitsProcessor(
371
+ block_from_id=ASR_BLOCK_TOKEN_ID_FROM,
372
+ block_token_ids=state.extra_block_token_ids,
373
+ )
374
+ ]
375
+ )
376
+
377
+ with torch.inference_mode():
378
+ outputs = state.model.generate(**inputs, **generate_kwargs)
379
+
380
+ input_ids = inputs["input_ids"]
381
+ generated_ids = outputs[0][len(input_ids[0].tolist()) :]
382
+ prediction_raw = state.tokenizer.decode(generated_ids, skip_special_tokens=False)
383
+ return normalize_prediction_text(prediction_raw), int(input_ids.shape[-1])
384
+
385
+
386
+ async def transcribe(
387
+ audio: str | tuple[int, Any] | None,
388
+ max_new_tokens: int,
389
+ begin_time: float,
390
+ end_time: float,
391
+ ) -> str:
392
+ started = time.perf_counter()
393
+ tmp_path: str | None = None
394
+ try:
395
+ logger.info("Transcribe request started")
396
+ await ensure_loaded()
397
+ audio_path, tmp_path = audio_to_path(audio)
398
+ async with infer_lock:
399
+ text, prompt_tokens = await asyncio.to_thread(
400
+ run_transformers_generation,
401
+ audio_path,
402
+ begin_time,
403
+ end_time,
404
+ int(max_new_tokens),
405
+ )
406
+ elapsed = time.perf_counter() - started
407
+ logger.info("Transcribe request finished in %.2fs", elapsed)
408
+ logger.info(
409
+ "Generation metadata: prompt_tokens=%s model=%s backend=transformers/%s %s",
410
+ prompt_tokens,
411
+ MODEL_ID,
412
+ state.resolved_attn_impl,
413
+ maybe_gpu_memory_text(),
414
+ )
415
+ return text
416
+ except gr.Error:
417
+ raise
418
+ except Exception as exc:
419
+ logger.exception("ASR request failed")
420
+ raise gr.Error(f"{exc.__class__.__name__}: {exc}") from exc
421
+ finally:
422
+ if tmp_path:
423
+ try:
424
+ os.unlink(tmp_path)
425
+ except OSError:
426
+ pass
427
+
428
+
429
+ APP_CSS = """
430
+ .gradio-container {
431
+ max-width: 1120px !important;
432
+ margin: 0 auto !important;
433
+ background:
434
+ radial-gradient(circle at top left, rgba(27, 99, 146, 0.12), transparent 34rem),
435
+ linear-gradient(180deg, #f6f8fb 0%, #eef3f7 100%);
436
+ color: #172033;
437
+ }
438
+ .ark-header {
439
+ padding: 26px 4px 20px;
440
+ border-bottom: 1px solid rgba(23, 32, 51, 0.12);
441
+ margin-bottom: 18px;
442
+ }
443
+ .ark-eyebrow {
444
+ margin: 0 0 7px;
445
+ color: #536579;
446
+ font-size: 15px;
447
+ font-weight: 700;
448
+ letter-spacing: 0.04em;
449
+ text-transform: uppercase;
450
+ }
451
+ .ark-title {
452
+ margin: 0;
453
+ color: #101828;
454
+ font-size: 40px;
455
+ line-height: 1.1;
456
+ font-weight: 800;
457
+ }
458
+ .ark-subtitle {
459
+ max-width: 920px;
460
+ margin: 12px 0 0;
461
+ color: #405166;
462
+ font-size: 17px;
463
+ line-height: 1.5;
464
+ }
465
+ .ark-opd {
466
+ color: #0b5cad;
467
+ font-weight: 800;
468
+ }
469
+ .ark-badges {
470
+ display: flex;
471
+ flex-wrap: wrap;
472
+ justify-content: flex-start;
473
+ gap: 8px;
474
+ margin-top: 14px;
475
+ }
476
+ .ark-badge {
477
+ display: inline-flex;
478
+ align-items: center;
479
+ height: 28px;
480
+ border-radius: 4px;
481
+ text-decoration: none !important;
482
+ background: transparent;
483
+ box-shadow: 0 1px 2px rgba(16, 24, 40, 0.12);
484
+ }
485
+ .ark-badge img {
486
+ display: block;
487
+ height: 28px;
488
+ }
489
+ .ark-panel {
490
+ border: 1px solid rgba(23, 32, 51, 0.12);
491
+ border-radius: 8px;
492
+ background: rgba(255, 255, 255, 0.9);
493
+ padding: 16px;
494
+ box-shadow: 0 10px 32px rgba(16, 24, 40, 0.06);
495
+ }
496
+ .ark-panel textarea {
497
+ font-size: 17px !important;
498
+ line-height: 1.65 !important;
499
+ }
500
+ .ark-panel button.primary,
501
+ .ark-panel button[variant="primary"] {
502
+ border-radius: 8px !important;
503
+ }
504
+ @media (max-width: 760px) {
505
+ .ark-badges {
506
+ max-width: 100%;
507
+ }
508
+ .ark-title {
509
+ font-size: 32px;
510
+ }
511
+ .ark-subtitle {
512
+ font-size: 16px;
513
+ }
514
+ }
515
+ """
516
+
517
+
518
+ with gr.Blocks(title="Ark ASR 0.6B", css=APP_CSS) as demo:
519
+ gr.HTML(
520
+ """
521
+ <header class="ark-header">
522
+ <p class="ark-eyebrow">Industrial Audio Online Policy Distillation</p>
523
+ <h1 class="ark-title">Ark ASR 0.6B</h1>
524
+ <p class="ark-subtitle"><span class="ark-opd">Open Audio OPD</span> brings online policy distillation to ASR, with the best overall results among the 0.6B-scale ASR models compared in the project.</p>
525
+ <nav class="ark-badges" aria-label="Project links">
526
+ <a class="ark-badge" href="https://github.com/AutoArk/open-audio-opd" target="_blank" rel="noopener noreferrer"><img src="https://img.shields.io/badge/GitHub-open--audio--opd-black?style=for-the-badge&logo=github&logoColor=white" alt="GitHub open-audio-opd"></a>
527
+ <a class="ark-badge" href="https://huggingface.co/AutoArk-AI/ARK-ASR-0.6B" target="_blank" rel="noopener noreferrer"><img src="https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-ARK--ASR--0.6B-yellow?style=for-the-badge" alt="Hugging Face ARK-ASR-0.6B"></a>
528
+ <a class="ark-badge" href="https://github.com/AutoArk/open-audio-opd/blob/main/paper/arxiv_ark_asr_opd/main.pdf" target="_blank" rel="noopener noreferrer"><img src="https://img.shields.io/badge/Paper-PDF-b31b1b?style=for-the-badge&logo=readthedocs&logoColor=white" alt="Paper PDF"></a>
529
+ <a class="ark-badge" href="https://github.com/AutoArk/open-audio-opd/blob/main/LICENSE" target="_blank" rel="noopener noreferrer"><img src="https://img.shields.io/badge/License-See%20LICENSE-blue?style=for-the-badge" alt="License"></a>
530
+ </nav>
531
+ </header>
532
+ """
533
+ )
534
+ with gr.Row(equal_height=True):
535
+ with gr.Column(scale=1, elem_classes=["ark-panel"]):
536
+ audio_input = gr.Audio(
537
+ sources=["upload", "microphone"],
538
+ type="filepath",
539
+ label="Audio",
540
+ )
541
+ with gr.Row():
542
+ begin_input = gr.Number(value=-1, label="Begin time")
543
+ end_input = gr.Number(value=-1, label="End time")
544
+ max_tokens_input = gr.Slider(
545
+ minimum=16,
546
+ maximum=512,
547
+ value=MAX_NEW_TOKENS,
548
+ step=16,
549
+ label="Max new tokens",
550
+ )
551
+ transcribe_button = gr.Button("Transcribe", variant="primary")
552
+ with gr.Column(scale=1, elem_classes=["ark-panel"]):
553
+ text_output = gr.Textbox(label="Transcript", lines=8)
554
+
555
+ transcribe_button.click(
556
+ transcribe,
557
+ inputs=[audio_input, max_tokens_input, begin_input, end_input],
558
+ outputs=text_output,
559
+ )
560
+
561
+
562
+ if __name__ == "__main__":
563
+ demo.queue(default_concurrency_limit=1).launch()
packages.txt ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ ffmpeg
2
+ libsndfile1
requirements.txt ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ --extra-index-url https://download.pytorch.org/whl/cu128
2
+ torch==2.9.0
3
+ torchaudio==2.9.0
4
+ transformers==4.57.3
5
+ gradio==5.50.0
6
+ huggingface-hub>=0.34.0
7
+ soundfile>=0.13.1
8
+ librosa>=0.11.0
9
+ numpy>=2.0
scripts/run_local_gradio.sh ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ set -euo pipefail
3
+
4
+ ROOT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)"
5
+
6
+ PYTHON_BIN="${PYTHON_BIN:-/root/miniforge3/envs/asr_vlm/bin/python}"
7
+ MODEL_PATH="${MODEL_PATH:-/data/yumu/model/trained_model/ark_asr_td_opd}"
8
+ HOST="${HOST:-0.0.0.0}"
9
+ PORT="${PORT:-18096}"
10
+ GPU="${GPU:-1}"
11
+ LOG_FILE="${LOG_FILE:-${ROOT_DIR}/runs/gradio_space.log}"
12
+ PID_FILE="${PID_FILE:-${ROOT_DIR}/runs/gradio_space.pid}"
13
+
14
+ mkdir -p "$(dirname "${LOG_FILE}")"
15
+
16
+ if [[ -s "${PID_FILE}" ]]; then
17
+ old_pid="$(cat "${PID_FILE}")"
18
+ if [[ -n "${old_pid}" ]] && kill -0 "${old_pid}" 2>/dev/null; then
19
+ kill "${old_pid}" 2>/dev/null || true
20
+ sleep 1
21
+ fi
22
+ fi
23
+
24
+ cd "${ROOT_DIR}"
25
+
26
+ CUDA_VISIBLE_DEVICES="${GPU}" \
27
+ ARK_ASR_MODEL_ID="${MODEL_PATH}" \
28
+ GRADIO_SERVER_NAME="${HOST}" \
29
+ GRADIO_SERVER_PORT="${PORT}" \
30
+ setsid "${PYTHON_BIN}" app.py > "${LOG_FILE}" 2>&1 < /dev/null &
31
+
32
+ echo "$!" > "${PID_FILE}"
33
+ echo "Started local Gradio: pid=$(cat "${PID_FILE}") url=http://${HOST}:${PORT}"
34
+ echo "Model: ${MODEL_PATH}"
35
+ echo "Log: ${LOG_FILE}"