CoolFace
Modelpublic

hypermind-official/ARK-ASR-3B-NoTranslate

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
1likes44downloads
processing_arkasr.py447 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 # 显式引入 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