Edge0/GPA-v1.5
2774
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 # Explicitly import soundfile to handle BytesIO.14 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 77 @classmethod78 def from_pretrained(cls, pretrained_model_name_or_path: str, **kwargs) -> "ArkasrProcessor":79 trust_remote_code = bool(kwargs.pop("trust_remote_code", False))80 passthrough_keys = {"cache_dir", "force_download", "local_files_only", "token", "revision", "subfolder"}81 shared_kwargs = {k: kwargs[k] for k in list(kwargs.keys()) if k in passthrough_keys}82 83 merge_factor = 484 audio_token = "<|audio|>"85 audio_dtype = "float32"86 tokenizer_cfg: Dict[str, Any] = {}87 feat_cfg: Dict[str, Any] = {}88 89 proc_cfg_path = os.path.join(pretrained_model_name_or_path, "processor_config.json")90 if os.path.isfile(proc_cfg_path):91 with open(proc_cfg_path, "r", encoding="utf-8") as f:92 proc_cfg = json.load(f)93 merge_factor = int(proc_cfg.get("merge_factor", merge_factor))94 audio_token = str(proc_cfg.get("audio_token", audio_token))95 audio_dtype = str(proc_cfg.get("audio_dtype", audio_dtype))96 tokenizer_cfg = proc_cfg.get("tokenizer_config", {}) or {}97 feat_cfg = proc_cfg.get("feature_extractor_config", {}) or {}98 99 feature_extractor = WhisperFeatureExtractor.from_pretrained(pretrained_model_name_or_path, **shared_kwargs)100 for k, v in feat_cfg.items():101 if hasattr(feature_extractor, k):102 try: setattr(feature_extractor, k, v)103 except Exception: pass104 105 tokenizer = AutoTokenizer.from_pretrained(106 pretrained_model_name_or_path, use_fast=True, trust_remote_code=trust_remote_code, **shared_kwargs107 )108 for k, v in tokenizer_cfg.items():109 if hasattr(tokenizer, k):110 try: setattr(tokenizer, k, v)111 except Exception: pass112 113 return cls(114 feature_extractor=feature_extractor,115 tokenizer=tokenizer,116 merge_factor=merge_factor,117 audio_token=audio_token,118 audio_dtype=audio_dtype,119 )120 121 # =========================122 # audio helpers (Modified)123 # =========================124 def _load_audio_file(self, path: str, sampling_rate: int = 16000, offset: float = 0.0, duration: Optional[float] = None) -> np.ndarray:125 # librosa.load supports offset and duration.126 # offset: start reading after this time (in seconds)127 # duration: only load up to this much audio (in seconds)128 audio_array, _ = librosa.load(path, sr=int(sampling_rate), mono=True, offset=offset, duration=duration)129 return np.asarray(audio_array, dtype=np.float32)130 131 def _strip_data_url_prefix(self, b64: str) -> str:132 if "," in b64 and b64[:30].lower().startswith("data:"):133 return b64.split(",", 1)[1]134 return b64135 136 def _load_audio_base64(self, b64: str, sampling_rate: int = 16000, offset: float = 0.0, duration: Optional[float] = None) -> np.ndarray:137 b64 = self._strip_data_url_prefix(b64)138 raw = base64.b64decode(b64)139 bio = io.BytesIO(raw)140 141 # librosa also supports offset and duration when loading from BytesIO.142 try:143 wav, _sr = librosa.load(bio, sr=int(sampling_rate), mono=True, offset=offset, duration=duration)144 return np.asarray(wav, dtype=np.float32)145 except Exception as e:146 # Fallback path: manual slicing, which is slower.147 try:148 bio.seek(0)149 data, sr = sf.read(bio, dtype="float32", always_2d=True)150 wav = data.mean(axis=1)151 if int(sr) != int(sampling_rate):152 wav = librosa.resample(wav, orig_sr=int(sr), target_sr=int(sampling_rate))153 154 start_sample = int(offset * sampling_rate)155 end_sample = None156 if duration is not None:157 end_sample = start_sample + int(duration * sampling_rate)158 159 return np.asarray(wav[start_sample:end_sample], dtype=np.float32)160 except Exception as e2:161 raise ValueError("Failed to decode base64 audio.") from e2162 163 def calculate_audio_token_count(self, mel_frames: int) -> int:164 downsampled = (int(mel_frames) + 1) // 2165 merged = downsampled // max(self.merge_factor, 1)166 return max(int(merged), 1)167 168 def _build_templates_and_audios(169 self,170 conversations: List[List[dict]],171 sampling_rate: int,172 add_generation_prompt: bool,173 ) -> tuple[List[str], List[np.ndarray], List[int]]:174 prompts_template: List[str] = []175 audios_raw: List[np.ndarray] = []176 prompt_audio_counts: List[int] = []177 178 for conv in conversations:179 conv_str = ""180 last_role = None181 audio_count_this_conv = 0182 183 for msg in conv:184 role = msg["role"]185 last_role = role186 content = msg["content"]187 188 if role == "user": conv_str += f"{self.user_token}"189 elif role == "assistant": conv_str += f"{self.assistant_token}"190 else: conv_str += f"<|{role}|>"191 192 if isinstance(content, str):193 conv_str += f"{content}"194 elif isinstance(content, list):195 for part in content:196 ptype = part.get("type")197 if ptype == "audio":198 # ------------------------------------------------------------199 # Parse begin_time and end_time when present.200 # ------------------------------------------------------------201 begin_time = part.get("begin_time", -1)202 end_time = part.get("end_time", -1)203 204 offset = 0.0205 duration = None206 207 # Apply slicing only when begin_time is valid and non-negative.208 if begin_time is not None and begin_time >= 0:209 offset = float(begin_time)210 if end_time is not None and end_time > begin_time:211 duration = float(end_time) - float(begin_time)212 213 audio_raw_this = None214 if "array" in part:215 arr = part["array"]216 if isinstance(arr, torch.Tensor):217 arr = arr.detach().cpu().numpy()218 full_arr = np.asarray(arr, dtype=np.float32).reshape(-1)219 220 # Slice the in-memory audio array.221 start_idx = int(offset * sampling_rate)222 end_idx = None223 if duration is not None:224 end_idx = start_idx + int(duration * sampling_rate)225 audio_raw_this = full_arr[start_idx:end_idx]226 227 elif "path" in part:228 audio_raw_this = self._load_audio_file(229 part["path"], 230 sampling_rate=sampling_rate,231 offset=offset,232 duration=duration233 )234 elif "base64" in part:235 audio_raw_this = self._load_audio_base64(236 part["base64"], 237 sampling_rate=sampling_rate,238 offset=offset,239 duration=duration240 )241 else:242 raise ValueError("Audio part must contain 'path' or 'array' or 'base64'.")243 244 audios_raw.append(audio_raw_this)245 audio_count_this_conv += 1246 conv_str += f"{self.bos_audio_token}{_AUDIO_MARKER}{self.eos_audio_token}"247 248 elif ptype == "text":249 conv_str += f"{part.get('text', '')}"250 else:251 raise ValueError(f"Unknown content part type: {ptype}")252 else:253 raise ValueError(f"Unsupported message content type: {type(content)}")254 255 if add_generation_prompt:256 if last_role == "user": conv_str += f"{self.assistant_token}"257 elif last_role == "assistant": conv_str += f"{self.user_token}"258 else: conv_str += f"{self.assistant_token}"259 260 prompts_template.append(conv_str)261 prompt_audio_counts.append(audio_count_this_conv)262 263 return prompts_template, audios_raw, prompt_audio_counts264 265 def _calculate_audio_token_counts_per_sample(266 self,267 audios_raw: List[np.ndarray],268 sampling_rate: int,269 audio_max_length: Optional[int],270 audio_pad_to_multiple_of: Optional[int],271 ) -> List[int]:272 del sampling_rate, audio_pad_to_multiple_of273 274 hop_length = int(getattr(self.feature_extractor, "hop_length", 160))275 max_audio_samples = int(audio_max_length) if audio_max_length is not None else None276 token_counts: List[int] = []277 278 for audio_raw in audios_raw:279 audio_np = np.asarray(audio_raw, dtype=np.float32).reshape(-1)280 effective_len = int(audio_np.shape[0])281 if max_audio_samples is not None:282 effective_len = min(effective_len, max_audio_samples)283 284 mel_frames = effective_len // max(hop_length, 1)285 token_counts.append(self.calculate_audio_token_count(int(mel_frames)))286 287 return token_counts288 289 # =========================290 # apply_chat_template291 # =========================292 def apply_chat_template(293 self,294 conversation: Union[List[dict], List[List[dict]]],295 chat_template: Optional[str] = None,296 add_generation_prompt: bool = True,297 **kwargs,298 ) -> Union[BatchFeature, str, List[str]]:299 if chat_template is not None:300 logger.warning("chat_template argument is ignored.")301 302 tokenize = kwargs.pop("tokenize", True)303 return_tensors = kwargs.pop("return_tensors", "pt")304 kwargs.pop("return_dict", None)305 306 audio_torch_dtype = kwargs.pop("audio_torch_dtype", None)307 audio_dtype_override = kwargs.pop("audio_dtype", None)308 dtype_source = audio_torch_dtype if audio_torch_dtype is not None else audio_dtype_override309 target_dtype = _resolve_torch_dtype(dtype_source, default=getattr(self, "audio_dtype", "float32"))310 311 text_kwargs = dict(kwargs.pop("text_kwargs", {}) or {})312 for k in ("padding", "truncation", "max_length", "add_special_tokens"):313 if k in kwargs and k not in text_kwargs:314 text_kwargs[k] = kwargs.pop(k)315 316 sampling_rate = int(kwargs.pop("sampling_rate", 16000))317 audio_padding = kwargs.pop("audio_padding", "longest")318 audio_max_length = kwargs.pop("audio_max_length", None)319 audio_pad_to_multiple_of = kwargs.pop("audio_pad_to_multiple_of", None)320 321 if kwargs:322 logger.warning(f"Ignored unused kwargs: {list(kwargs.keys())}")323 324 if isinstance(conversation, list) and conversation and isinstance(conversation[0], dict):325 conversations = [conversation]326 is_single = True327 else:328 conversations = conversation329 is_single = False330 331 prompt_templates, audios_raw, prompt_audio_counts = self._build_templates_and_audios(332 conversations=conversations,333 sampling_rate=sampling_rate,334 add_generation_prompt=add_generation_prompt,335 )336 337 input_features = None338 audio_token_counts: List[int] = []339 340 if len(audios_raw) > 0:341 feat = self.feature_extractor(342 audios_raw,343 sampling_rate=sampling_rate,344 return_tensors="np",345 return_attention_mask=False,346 padding=audio_padding,347 max_length=audio_max_length,348 pad_to_multiple_of=audio_pad_to_multiple_of,349 )350 input_features = feat["input_features"]351 if not isinstance(input_features, np.ndarray):352 input_features = np.asarray(input_features)353 354 audio_token_counts = self._calculate_audio_token_counts_per_sample(355 audios_raw=audios_raw,356 sampling_rate=sampling_rate,357 audio_max_length=audio_max_length,358 audio_pad_to_multiple_of=audio_pad_to_multiple_of,359 )360 361 prompts: List[str] = []362 audio_idx = 0363 for prompt_template, audio_count in zip(prompt_templates, prompt_audio_counts):364 prompt = prompt_template365 for _ in range(audio_count):366 if audio_idx >= len(audio_token_counts):367 raise ValueError("Audio token count mismatch while building prompts.")368 audio_tokens_str = "".join([self.audio_token] * audio_token_counts[audio_idx])369 prompt = prompt.replace(_AUDIO_MARKER, audio_tokens_str, 1)370 audio_idx += 1371 if _AUDIO_MARKER in prompt:372 raise ValueError("Unresolved audio marker remained in prompt.")373 prompts.append(prompt)374 375 if audio_idx != len(audio_token_counts):376 raise ValueError("Unused audio token counts remained after prompt construction.")377 378 if not tokenize:379 return prompts[0] if is_single else prompts380 381 text_kwargs.setdefault("padding", "longest")382 text_kwargs.setdefault("add_special_tokens", False)383 text_kwargs["return_tensors"] = return_tensors384 385 enc = self.tokenizer(prompts, **text_kwargs)386 data: Dict[str, Any] = dict(enc)387 388 if input_features is not None:389 data["audios"] = torch.tensor(input_features, dtype=target_dtype)390 391 return BatchFeature(data=data, tensor_type=return_tensors)392 393 # ... (The remaining batch_decode, decode, __call__, and model_input_names stay unchanged.) ...394 def batch_decode(self, *args, **kwargs):395 return self.tokenizer.batch_decode(*args, **kwargs)396 397 def decode(self, *args, **kwargs):398 return self.tokenizer.decode(*args, **kwargs)399 400 def __call__(401 self,402 text: Union[str, List[str]],403 audios: Union[np.ndarray, torch.Tensor, List[Union[np.ndarray, torch.Tensor]]],404 sampling_rate: int = 16000,405 return_tensors: str = "pt",406 **tokenizer_kwargs,407 ) -> BatchFeature:408 # Simplified implementation that skips time slicing because the caller passes raw audio arrays directly.409 audios_list = []410 def flatten_audios(obj):411 if isinstance(obj, (list, tuple)):412 if len(obj) > 0 and isinstance(obj[0], (float, int)):413 audios_list.append(obj)414 else:415 for item in obj: flatten_audios(item)416 elif isinstance(obj, (np.ndarray, torch.Tensor)):417 audios_list.append(obj)418 flatten_audios(audios)419 420 audios_np: List[np.ndarray] = []421 for a in audios_list:422 if isinstance(a, torch.Tensor): a = a.detach().cpu().numpy()423 a = np.asarray(a, dtype=np.float32).reshape(-1)424 audios_np.append(a)425 426 input_features = None427 if audios_np:428 feat = self.feature_extractor(audios_np, sampling_rate=int(sampling_rate), return_tensors="np", return_attention_mask=False, padding="longest")429 input_features = feat["input_features"]430 if not isinstance(input_features, np.ndarray): input_features = np.asarray(input_features)431 432 tokenizer_kwargs = dict(tokenizer_kwargs or {})433 tokenizer_kwargs.setdefault("padding", "longest")434 tokenizer_kwargs.setdefault("add_special_tokens", False)435 tokenizer_kwargs["return_tensors"] = return_tensors436 437 enc = self.tokenizer(text, **tokenizer_kwargs)438 data: Dict[str, Any] = dict(enc)439 if input_features is not None:440 data["audios"] = torch.tensor(input_features, dtype=_resolve_torch_dtype(getattr(self, "audio_dtype", "float32")))441 return BatchFeature(data=data, tensor_type=return_tensors)442 443 @property444 def model_input_names(self):445 return ["input_ids", "attention_mask", "audios"]446 