CoolFace
Modelpublic

Edge0/GPA-v1.5

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
27likes74downloads
processing_arkasr.py446 linesDownload Raw Back to root
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