EdgeIn-v3 / processing_edgeinstant.py
chenjz24's picture
Upload folder using huggingface_hub
d81a465 verified
Raw History Blame Contribute Delete
13 kB
"""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<think>\n"
return prefix if enable_thinking else prefix + "\n</think>\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()