"""Text and 16 kHz speech inputs for EdgeInstant.""" from __future__ import annotations import json from pathlib import Path import numpy as np from transformers import AutoImageProcessor, AutoTokenizer, ProcessorMixin, WhisperFeatureExtractor from transformers.dynamic_module_utils import custom_object_save from transformers.feature_extraction_utils import BatchFeature from .configuration_edgeinstant import EdgeInstantConfig class EdgeInstantProcessor(ProcessorMixin): tokenizer_class = "AutoTokenizer" feature_extractor_class = "WhisperFeatureExtractor" model_input_names = ["input_ids", "attention_mask", "input_features", "feature_attention_mask"] def __init__(self, tokenizer, feature_extractor, config, image_processor=None, video_processor_config=None): self.tokenizer = tokenizer self.feature_extractor = feature_extractor self.config = config self.image_processor = image_processor self.video_processor_config = video_processor_config self.chat_template = getattr(tokenizer, "chat_template", None) def audio_token_count(self, frame_count): """Count encoder frames, temporal grouping and projector output tokens.""" frame_count = int(frame_count) # Qwen3-ASR uses three stride-2 convolutions per 100-frame block. count = (frame_count // 100) * 13 + ((frame_count % 100) + 7) // 8 stride = self.config.audio_stack if self.config.audio_stack > 1 else self.config.audio_pool count = (count + stride - 1) // stride projector = self.config.projector_config if projector["align_mode"] == "salmonn_qformer": window = int(projector["window"]) return (count + window - 1) // window if projector["align_mode"] != "residual": hidden_size = self.config.thinker_config.text_config.hidden_size count = count * projector["output_size"] // hidden_size return count def _audio_batch(self, audio, sampling_rate): if sampling_rate != self.config.sampling_rate: raise ValueError(f"Audio must use {self.config.sampling_rate} Hz; received {sampling_rate} Hz") if hasattr(audio, "detach"): audio = audio.detach().cpu().numpy() if isinstance(audio, np.ndarray): batch = [audio] if audio.ndim == 1 else list(audio) elif isinstance(audio, (list, tuple)): batch = [audio] if audio and np.isscalar(audio[0]) else list(audio) else: raise TypeError("audio must be a mono waveform or a batch of mono waveforms") if not batch: raise ValueError("Audio must be a nonempty mono waveform") result = [] for waveform in batch: if hasattr(waveform, "detach"): waveform = waveform.detach().cpu().numpy() waveform = np.asarray(waveform, dtype=np.float32) if waveform.ndim != 1 or not waveform.size: raise ValueError("Audio must be a nonempty mono waveform") result.append(waveform) return result def audio_features(self, audio, sampling_rate=16000, return_tensors="np"): """Extract complete waveforms with zero context for their final mel frames.""" waveforms = self._audio_batch(audio, sampling_rate) hop = self.feature_extractor.hop_length longest = max(len(waveform) for waveform in waveforms) aligned_samples = ((longest + hop - 1) // hop) * hop original_padding = max(self.feature_extractor.n_samples, aligned_samples) context_padding = ((aligned_samples + self.feature_extractor.n_fft + hop - 1) // hop) * hop # Preserve the original STFT reflection boundary for full-length recordings. padded_samples = min(original_padding, context_padding) return self.feature_extractor( waveforms, sampling_rate=self.config.sampling_rate, padding="max_length", max_length=int(padded_samples), truncation=False, return_attention_mask=True, return_tensors=return_tensors, ) @staticmethod def assistant_prefix(enable_thinking=False): prefix = "<|im_start|>assistant\n\n" return prefix if enable_thinking else prefix + "\n\n\n" @staticmethod def _batch_flags(value, batch_size, name): flags = [value] * batch_size if isinstance(value, bool) else list(value) if len(flags) != batch_size or any(not isinstance(flag, bool) for flag in flags): raise ValueError(f"{name} must be a bool or one bool per input") return flags def __call__( self, text=None, audio=None, sampling_rate=16000, task="qa", system_prompt=None, history=None, omit_audio_prompt=False, enable_thinking=False, audio_first=True, return_tensors="pt", padding=True, **kwargs, ): """Build a generation prompt for each text or speech input. A string prompt is shared by all waveforms; a list supplies one prompt per waveform. ``enable_thinking`` accepts a bool or one bool per input. ``audio_first`` controls whether audio precedes the user text, with the same scalar or per-input format. ``history`` supplies preceding text turns shared by the batch. ``omit_audio_prompt`` renders a user turn containing only the audio. Text padding defaults to the left for batched generation; pass ``padding_side="right"`` for training batches. """ if task not in {"qa", "asr"}: raise ValueError(f"Unsupported audio task: {task}") if text is None and audio is None: raise ValueError("Provide text or audio") history = history or [] if any(message.get("role") not in {"user", "assistant"} or not isinstance(message.get("content"), str) for message in history): raise ValueError("history requires user/assistant text messages") if omit_audio_prompt and (audio is None or task != "qa" or text): raise ValueError("omit_audio_prompt requires audio, task=qa, and no text prompt") history_prefix = "".join( f"<|im_start|>{message['role']}\n{message['content']}<|im_end|>\n" for message in history ) audio_features = {} if audio is not None: waveforms = self._audio_batch(audio, sampling_rate) features = self.audio_features(waveforms, sampling_rate=sampling_rate) audio_features = { "input_features": features["input_features"], "feature_attention_mask": features["attention_mask"], } counts = [self.audio_token_count(length) for length in features["attention_mask"].sum(-1)] prompts = [text or ""] * len(waveforms) if text is None or isinstance(text, str) else list(text) if len(prompts) != len(waveforms): raise ValueError("Provide one text prompt per waveform, or one shared string prompt") thinking_flags = self._batch_flags(enable_thinking, len(prompts), "enable_thinking") audio_first_flags = self._batch_flags(audio_first, len(prompts), "audio_first") rendered = [] for prompt, count, thinking, first in zip(prompts, counts, thinking_flags, audio_first_flags): user_prompt = prompt.strip() or ("Transcribe the speech." if task == "asr" else self.config.qa_prompt) prefix = f"<|im_start|>system\n{system_prompt}<|im_end|>\n" if system_prompt and task != "asr" else "" prefix += history_prefix audio_prompt = f"<|audio_start|>{'<|audio_pad|>' * count}<|audio_end|>" user_content = f"{audio_prompt}\n{user_prompt}" if first else f"{user_prompt}\n{audio_prompt}" if omit_audio_prompt: user_content = audio_prompt rendered.append( f"{prefix}<|im_start|>user\n{user_content}<|im_end|>\n{self.assistant_prefix(thinking)}" ) else: prompts = [text] if isinstance(text, str) else list(text) thinking_flags = self._batch_flags(enable_thinking, len(prompts), "enable_thinking") rendered = [] for prompt, thinking in zip(prompts, thinking_flags): messages = [] if system_prompt: messages.append({"role": "system", "content": system_prompt}) messages.extend(history) messages.append({"role": "user", "content": prompt}) rendered.append(self.tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, enable_thinking=thinking, )) kwargs.setdefault("padding_side", "left") kwargs.setdefault("return_token_type_ids", False) tokens = self.tokenizer( rendered, add_special_tokens=False, padding=padding, return_tensors=return_tensors, **kwargs, ) return BatchFeature(data={**tokens, **audio_features}, tensor_type=return_tensors) def decode(self, *args, **kwargs): return self.tokenizer.decode(*args, **kwargs) def batch_decode(self, *args, **kwargs): return self.tokenizer.batch_decode(*args, **kwargs) def to_dict(self): return { "processor_class": self.__class__.__name__, "auto_map": {"AutoProcessor": "processing_edgeinstant.EdgeInstantProcessor"}, "image_processor_subfolder": "image_processor" if self.image_processor is not None else None, } def save_pretrained(self, save_directory, **kwargs): directory = Path(save_directory) directory.mkdir(parents=True, exist_ok=True) self.config.save_pretrained(directory) self.tokenizer.save_pretrained(directory, **kwargs) self.feature_extractor.save_pretrained(directory) if self.image_processor is not None: self.image_processor.save_pretrained(directory / "image_processor") if self.video_processor_config is not None: (directory / "image_processor" / "video_preprocessor_config.json").write_text( json.dumps(self.video_processor_config, indent=2) + "\n", encoding="utf-8", ) processor_file = directory / "processor_config.json" processor_file.write_text(json.dumps(self.to_dict(), indent=2) + "\n", encoding="utf-8") custom_object_save(self, directory) return [str(processor_file)] @classmethod def from_pretrained(cls, pretrained_model_name_or_path, **kwargs): trust_remote_code = kwargs.pop("trust_remote_code", False) kwargs.pop("_from_auto", None) config = kwargs.pop("config", None) hub_keys = { "cache_dir", "force_download", "local_files_only", "token", "revision", "subfolder", "proxies", } hub_kwargs = {key: value for key, value in kwargs.items() if key in hub_keys} if config is None: config = EdgeInstantConfig.from_pretrained(pretrained_model_name_or_path, **hub_kwargs) tokenizer = AutoTokenizer.from_pretrained( pretrained_model_name_or_path, config=config.thinker_config, trust_remote_code=trust_remote_code, **kwargs, ) feature_extractor = WhisperFeatureExtractor.from_pretrained(pretrained_model_name_or_path, **hub_kwargs) metadata_path = Path(pretrained_model_name_or_path) / "processor_config.json" if not metadata_path.is_file(): from transformers.utils.hub import cached_file metadata_path = Path(cached_file(pretrained_model_name_or_path, "processor_config.json", **hub_kwargs)) metadata = json.loads(metadata_path.read_text()) image_processor = None video_processor_config = None if metadata.get("image_processor_subfolder"): subfolder = str(Path(hub_kwargs.get("subfolder", "")) / metadata["image_processor_subfolder"]) image_kwargs = {**hub_kwargs, "subfolder": subfolder} image_processor = AutoImageProcessor.from_pretrained(pretrained_model_name_or_path, **image_kwargs) from transformers.utils.hub import cached_file video_path = cached_file(pretrained_model_name_or_path, "video_preprocessor_config.json", _raise_exceptions_for_missing_entries=False, **image_kwargs) if video_path is not None: video_processor_config = json.loads(Path(video_path).read_text()) return cls(tokenizer=tokenizer, feature_extractor=feature_extractor, config=config, image_processor=image_processor, video_processor_config=video_processor_config) EdgeInstantProcessor.register_for_auto_class()