hypermind-official/ARK-ASR-3B-NoTranslate
144
1# coding=utf-82from __future__ import annotations3 4import base645import io6import json7import os8from typing import Any, Dict, List, Optional, Union9 10import numpy as np11import torch12import librosa13import soundfile as sf # 显式引入 soundfile 以处理 BytesIO14 15from transformers import AutoTokenizer, WhisperFeatureExtractor16from transformers.feature_extraction_utils import BatchFeature17from transformers.processing_utils import ProcessorMixin18from transformers.utils import logging19 20logger = logging.get_logger(__name__)21 22_AUDIO_MARKER = "<<AUDIO_TOKENS>>"23 24def _normalize_dtype_name(name: str) -> str:25 name = name.strip().lower()26 alias = {27 "fp16": "float16",28 "float16": "float16",29 "half": "float16",30 "bf16": "bfloat16",31 "bfloat16": "bfloat16",32 "fp32": "float32",33 "float32": "float32",34 "float": "float32",35 }36 return alias.get(name, name)37 38 39def _resolve_torch_dtype(x: Any, default: str = "float32") -> torch.dtype:40 if isinstance(x, torch.dtype):41 return x42 if x is None:43 x = default44 if isinstance(x, str):45 name = _normalize_dtype_name(x)46 if not hasattr(torch, name):47 raise ValueError(f"Unknown torch dtype string: {x} (normalized: {name})")48 return getattr(torch, name)49 raise TypeError(f"audio_dtype/audio_torch_dtype must be str or torch.dtype or None, got {type(x)}")50 51 52class ArkasrProcessor(ProcessorMixin):53 attributes = ["feature_extractor", "tokenizer"]54 valid_kwargs = ["merge_factor", "audio_token", "audio_dtype"]55 feature_extractor_class = ("WhisperFeatureExtractor", "SequenceFeatureExtractor")56 tokenizer_class = ("PreTrainedTokenizerFast", "PreTrainedTokenizer")57 58 def __init__(59 self,60 feature_extractor,61 tokenizer,62 merge_factor: int = 4,63 audio_token: str = "<|audio|>",64 audio_dtype: str = "float32",65 **kwargs,66 ):67 super().__init__(feature_extractor, tokenizer)68 self.merge_factor = int(merge_factor)69 self.audio_token = str(audio_token)70 self.audio_dtype = str(audio_dtype)71 72 self.bos_audio_token = "<|begin_of_audio|>"73 self.eos_audio_token = "<|end_of_audio|>"74 self.user_token = "<|user|>"75 self.assistant_token = "<|assistant|>"76 self.assistant_end_token = "<|im_end|>"77 78 @classmethod79 def from_pretrained(cls, pretrained_model_name_or_path: str, **kwargs) -> "ArkasrProcessor":80 trust_remote_code = bool(kwargs.pop("trust_remote_code", False))81 passthrough_keys = {"cache_dir", "force_download", "local_files_only", "token", "revision", "subfolder"}82 shared_kwargs = {k: kwargs[k] for k in list(kwargs.keys()) if k in passthrough_keys}83 84 merge_factor = 485 audio_token = "<|audio|>"86 audio_dtype = "float32"87 tokenizer_cfg: Dict[str, Any] = {}88 feat_cfg: Dict[str, Any] = {}89 90 proc_cfg_path = os.path.join(pretrained_model_name_or_path, "processor_config.json")91 if os.path.isfile(proc_cfg_path):92 with open(proc_cfg_path, "r", encoding="utf-8") as f:93 proc_cfg = json.load(f)94 merge_factor = int(proc_cfg.get("merge_factor", merge_factor))95 audio_token = str(proc_cfg.get("audio_token", audio_token))96 audio_dtype = str(proc_cfg.get("audio_dtype", audio_dtype))97 tokenizer_cfg = proc_cfg.get("tokenizer_config", {}) or {}98 feat_cfg = proc_cfg.get("feature_extractor_config", {}) or {}99 100 feature_extractor = WhisperFeatureExtractor.from_pretrained(pretrained_model_name_or_path, **shared_kwargs)101 for k, v in feat_cfg.items():102 if hasattr(feature_extractor, k):103 try: setattr(feature_extractor, k, v)104 except Exception: pass105 106 tokenizer = AutoTokenizer.from_pretrained(107 pretrained_model_name_or_path, use_fast=True, trust_remote_code=trust_remote_code, **shared_kwargs108 )109 for k, v in tokenizer_cfg.items():110 if hasattr(tokenizer, k):111 try: setattr(tokenizer, k, v)112 except Exception: pass113 114 return cls(115 feature_extractor=feature_extractor,116 tokenizer=tokenizer,117 merge_factor=merge_factor,118 audio_token=audio_token,119 audio_dtype=audio_dtype,120 )121 122 # =========================123 # audio helpers (Modified)124 # =========================125 def _load_audio_file(self, path: str, sampling_rate: int = 16000, offset: float = 0.0, duration: Optional[float] = None) -> np.ndarray:126 # librosa load 支持 offset 和 duration127 # offset: start reading after this time (in seconds)128 # duration: only load up to this much audio (in seconds)129 audio_array, _ = librosa.load(path, sr=int(sampling_rate), mono=True, offset=offset, duration=duration)130 return np.asarray(audio_array, dtype=np.float32)131 132 def _strip_data_url_prefix(self, b64: str) -> str:133 if "," in b64 and b64[:30].lower().startswith("data:"):134 return b64.split(",", 1)[1]135 return b64136 137 def _load_audio_base64(self, b64: str, sampling_rate: int = 16000, offset: float = 0.0, duration: Optional[float] = None) -> np.ndarray:138 b64 = self._strip_data_url_prefix(b64)139 raw = base64.b64decode(b64)140 bio = io.BytesIO(raw)141 142 # 使用 librosa 加载 BytesIO 同样支持 offset 和 duration143 try:144 wav, _sr = librosa.load(bio, sr=int(sampling_rate), mono=True, offset=offset, duration=duration)145 return np.asarray(wav, dtype=np.float32)146 except Exception as e:147 # Fallback (手动切片,比较慢)148 try:149 bio.seek(0)150 data, sr = sf.read(bio, dtype="float32", always_2d=True)151 wav = data.mean(axis=1)152 if int(sr) != int(sampling_rate):153 wav = librosa.resample(wav, orig_sr=int(sr), target_sr=int(sampling_rate))154 155 start_sample = int(offset * sampling_rate)156 end_sample = None157 if duration is not None:158 end_sample = start_sample + int(duration * sampling_rate)159 160 return np.asarray(wav[start_sample:end_sample], dtype=np.float32)161 except Exception as e2:162 raise ValueError("Failed to decode base64 audio.") from e2163 164 def calculate_audio_token_count(self, mel_frames: int) -> int:165 downsampled = (int(mel_frames) + 1) // 2166 merged = downsampled // max(self.merge_factor, 1)167 return max(int(merged), 1)168 169 def _build_templates_and_audios(170 self,171 conversations: List[List[dict]],172 sampling_rate: int,173 add_generation_prompt: bool,174 ) -> tuple[List[str], List[np.ndarray], List[int]]:175 prompts_template: List[str] = []176 audios_raw: List[np.ndarray] = []177 prompt_audio_counts: List[int] = []178 179 for conv in conversations:180 conv_str = ""181 last_role = None182 audio_count_this_conv = 0183 184 for msg in conv:185 role = msg["role"]186 last_role = role187 content = msg["content"]188 189 if role == "user": conv_str += f"{self.user_token}"190 elif role == "assistant": conv_str += f"{self.assistant_token}"191 else: conv_str += f"<|{role}|>"192 193 if isinstance(content, str):194 conv_str += f"{content}"195 elif isinstance(content, list):196 for part in content:197 ptype = part.get("type")198 if ptype == "audio":199 # ------------------------------------------------------------200 # 修改点:解析 begin_time 和 end_time201 # ------------------------------------------------------------202 begin_time = part.get("begin_time", -1)203 end_time = part.get("end_time", -1)204 205 offset = 0.0206 duration = None207 208 # 只有当 begin_time >= 0 且有效时才应用切片209 if begin_time is not None and begin_time >= 0:210 offset = float(begin_time)211 if end_time is not None and end_time > begin_time:212 duration = float(end_time) - float(begin_time)213 214 audio_raw_this = None215 if "array" in part:216 arr = part["array"]217 if isinstance(arr, torch.Tensor):218 arr = arr.detach().cpu().numpy()219 full_arr = np.asarray(arr, dtype=np.float32).reshape(-1)220 221 # 针对 array 的切片222 start_idx = int(offset * sampling_rate)223 end_idx = None224 if duration is not None:225 end_idx = start_idx + int(duration * sampling_rate)226 audio_raw_this = full_arr[start_idx:end_idx]227 228 elif "path" in part:229 audio_raw_this = self._load_audio_file(230 part["path"], 231 sampling_rate=sampling_rate,232 offset=offset,233 duration=duration234 )235 elif "base64" in part:236 audio_raw_this = self._load_audio_base64(237 part["base64"], 238 sampling_rate=sampling_rate,239 offset=offset,240 duration=duration241 )242 else:243 raise ValueError("Audio part must contain 'path' or 'array' or 'base64'.")244 245 audios_raw.append(audio_raw_this)246 audio_count_this_conv += 1247 conv_str += f"{self.bos_audio_token}{_AUDIO_MARKER}{self.eos_audio_token}"248 249 elif ptype == "text":250 conv_str += f"{part.get('text', '')}"251 else:252 raise ValueError(f"Unknown content part type: {ptype}")253 else:254 raise ValueError(f"Unsupported message content type: {type(content)}")255 256 if add_generation_prompt:257 if last_role == "user": conv_str += f"{self.assistant_token}"258 elif last_role == "assistant": conv_str += f"{self.assistant_end_token}"259 else: conv_str += f"{self.assistant_token}"260 261 prompts_template.append(conv_str)262 prompt_audio_counts.append(audio_count_this_conv)263 264 return prompts_template, audios_raw, prompt_audio_counts265 266 def _calculate_audio_token_counts_per_sample(267 self,268 audios_raw: List[np.ndarray],269 sampling_rate: int,270 audio_max_length: Optional[int],271 audio_pad_to_multiple_of: Optional[int],272 ) -> List[int]:273 del sampling_rate, audio_pad_to_multiple_of274 275 hop_length = int(getattr(self.feature_extractor, "hop_length", 160))276 max_audio_samples = int(audio_max_length) if audio_max_length is not None else None277 token_counts: List[int] = []278 279 for audio_raw in audios_raw:280 audio_np = np.asarray(audio_raw, dtype=np.float32).reshape(-1)281 effective_len = int(audio_np.shape[0])282 if max_audio_samples is not None:283 effective_len = min(effective_len, max_audio_samples)284 285 mel_frames = effective_len // max(hop_length, 1)286 token_counts.append(self.calculate_audio_token_count(int(mel_frames)))287 288 return token_counts289 290 # =========================291 # apply_chat_template292 # =========================293 def apply_chat_template(294 self,295 conversation: Union[List[dict], List[List[dict]]],296 chat_template: Optional[str] = None,297 add_generation_prompt: bool = True,298 **kwargs,299 ) -> Union[BatchFeature, str, List[str]]:300 if chat_template is not None:301 logger.warning("chat_template argument is ignored.")302 303 tokenize = kwargs.pop("tokenize", True)304 return_tensors = kwargs.pop("return_tensors", "pt")305 kwargs.pop("return_dict", None)306 307 audio_torch_dtype = kwargs.pop("audio_torch_dtype", None)308 audio_dtype_override = kwargs.pop("audio_dtype", None)309 dtype_source = audio_torch_dtype if audio_torch_dtype is not None else audio_dtype_override310 target_dtype = _resolve_torch_dtype(dtype_source, default=getattr(self, "audio_dtype", "float32"))311 312 text_kwargs = dict(kwargs.pop("text_kwargs", {}) or {})313 for k in ("padding", "truncation", "max_length", "add_special_tokens"):314 if k in kwargs and k not in text_kwargs:315 text_kwargs[k] = kwargs.pop(k)316 317 sampling_rate = int(kwargs.pop("sampling_rate", 16000))318 audio_padding = kwargs.pop("audio_padding", "longest")319 audio_max_length = kwargs.pop("audio_max_length", None)320 audio_pad_to_multiple_of = kwargs.pop("audio_pad_to_multiple_of", None)321 322 if kwargs:323 logger.warning(f"Ignored unused kwargs: {list(kwargs.keys())}")324 325 if isinstance(conversation, list) and conversation and isinstance(conversation[0], dict):326 conversations = [conversation]327 is_single = True328 else:329 conversations = conversation330 is_single = False331 332 prompt_templates, audios_raw, prompt_audio_counts = self._build_templates_and_audios(333 conversations=conversations,334 sampling_rate=sampling_rate,335 add_generation_prompt=add_generation_prompt,336 )337 338 input_features = None339 audio_token_counts: List[int] = []340 341 if len(audios_raw) > 0:342 feat = self.feature_extractor(343 audios_raw,344 sampling_rate=sampling_rate,345 return_tensors="np",346 return_attention_mask=False,347 padding=audio_padding,348 max_length=audio_max_length,349 pad_to_multiple_of=audio_pad_to_multiple_of,350 )351 input_features = feat["input_features"]352 if not isinstance(input_features, np.ndarray):353 input_features = np.asarray(input_features)354 355 audio_token_counts = self._calculate_audio_token_counts_per_sample(356 audios_raw=audios_raw,357 sampling_rate=sampling_rate,358 audio_max_length=audio_max_length,359 audio_pad_to_multiple_of=audio_pad_to_multiple_of,360 )361 362 prompts: List[str] = []363 audio_idx = 0364 for prompt_template, audio_count in zip(prompt_templates, prompt_audio_counts):365 prompt = prompt_template366 for _ in range(audio_count):367 if audio_idx >= len(audio_token_counts):368 raise ValueError("Audio token count mismatch while building prompts.")369 audio_tokens_str = "".join([self.audio_token] * audio_token_counts[audio_idx])370 prompt = prompt.replace(_AUDIO_MARKER, audio_tokens_str, 1)371 audio_idx += 1372 if _AUDIO_MARKER in prompt:373 raise ValueError("Unresolved audio marker remained in prompt.")374 prompts.append(prompt)375 376 if audio_idx != len(audio_token_counts):377 raise ValueError("Unused audio token counts remained after prompt construction.")378 379 if not tokenize:380 return prompts[0] if is_single else prompts381 382 text_kwargs.setdefault("padding", "longest")383 text_kwargs.setdefault("add_special_tokens", False)384 text_kwargs["return_tensors"] = return_tensors385 386 enc = self.tokenizer(prompts, **text_kwargs)387 data: Dict[str, Any] = dict(enc)388 389 if input_features is not None:390 data["audios"] = torch.tensor(input_features, dtype=target_dtype)391 392 return BatchFeature(data=data, tensor_type=return_tensors)393 394 # ... (其余 batch_decode, decode, __call__, model_input_names 保持不变) ...395 def batch_decode(self, *args, **kwargs):396 return self.tokenizer.batch_decode(*args, **kwargs)397 398 def decode(self, *args, **kwargs):399 return self.tokenizer.decode(*args, **kwargs)400 401 def __call__(402 self,403 text: Union[str, List[str]],404 audios: Union[np.ndarray, torch.Tensor, List[Union[np.ndarray, torch.Tensor]]],405 sampling_rate: int = 16000,406 return_tensors: str = "pt",407 **tokenizer_kwargs,408 ) -> BatchFeature:409 # 简化版实现,不包含时间切片逻辑,因为直接传入的是 audio array410 audios_list = []411 def flatten_audios(obj):412 if isinstance(obj, (list, tuple)):413 if len(obj) > 0 and isinstance(obj[0], (float, int)):414 audios_list.append(obj)415 else:416 for item in obj: flatten_audios(item)417 elif isinstance(obj, (np.ndarray, torch.Tensor)):418 audios_list.append(obj)419 flatten_audios(audios)420 421 audios_np: List[np.ndarray] = []422 for a in audios_list:423 if isinstance(a, torch.Tensor): a = a.detach().cpu().numpy()424 a = np.asarray(a, dtype=np.float32).reshape(-1)425 audios_np.append(a)426 427 input_features = None428 if audios_np:429 feat = self.feature_extractor(audios_np, sampling_rate=int(sampling_rate), return_tensors="np", return_attention_mask=False, padding="longest")430 input_features = feat["input_features"]431 if not isinstance(input_features, np.ndarray): input_features = np.asarray(input_features)432 433 tokenizer_kwargs = dict(tokenizer_kwargs or {})434 tokenizer_kwargs.setdefault("padding", "longest")435 tokenizer_kwargs.setdefault("add_special_tokens", False)436 tokenizer_kwargs["return_tensors"] = return_tensors437 438 enc = self.tokenizer(text, **tokenizer_kwargs)439 data: Dict[str, Any] = dict(enc)440 if input_features is not None:441 data["audios"] = torch.tensor(input_features, dtype=_resolve_torch_dtype(getattr(self, "audio_dtype", "float32")))442 return BatchFeature(data=data, tensor_type=return_tensors)443 444 @property445 def model_input_names(self):446 return ["input_ids", "attention_mask", "audios"]447 