Download inference.py from laion/humaneness-ears-base-medium: direct link, hf CLI and curl.
- Browser
- Download file 8.42 kB
-
https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/inference.py
- Command line
-
hf download hf://laion/humaneness-ears-base-medium/inference.py
-
curl -L -o inference.py https://huggingface.co/laion/humaneness-ears-base-medium/resolve/main/inference.py
8.42 kB
| #!/usr/bin/env python3 | |
| """Inference for LAION Humaneness Ears Base and Medium audio encoders.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import math | |
| from pathlib import Path | |
| import numpy as np | |
| import soundfile as sf | |
| import torch | |
| from safetensors.torch import load_file | |
| from scipy.signal import resample_poly | |
| from transformers import WhisperFeatureExtractor | |
| from model import LayeredMultiTaskWhisper | |
| def read_audio(path): | |
| wave, rate = sf.read(path, dtype='float32', always_2d=True) | |
| wave = wave.mean(axis=1) | |
| duration = len(wave) / rate | |
| if not .1 <= duration <= 30.: | |
| raise ValueError(f'Audio must last 0.1–30 seconds; received {duration:.3f}s') | |
| if not np.isfinite(wave).all(): | |
| raise ValueError('Audio contains nonfinite values') | |
| if rate != 16000: | |
| divisor = math.gcd(rate, 16000) | |
| wave = resample_poly(wave, 16000 // divisor, rate // divisor).astype(np.float32) | |
| return wave[:480000], duration | |
| class Predictor: | |
| """Load once and reuse. Local safe weights; no teachers or ASR decoder.""" | |
| def __init__(self, repo_dir, model_size='medium', device='cpu', threads=4): | |
| if model_size not in ('base', 'medium', 'small') or threads < 1: | |
| raise ValueError('Use base/medium and at least one CPU thread') | |
| # Accept the historical architecture label for existing callers. | |
| model_size = 'medium' if model_size == 'small' else model_size | |
| torch.set_num_threads(threads) | |
| self.repo = Path(repo_dir) | |
| self.model_size, self.device = model_size, device | |
| self.norm = json.loads((self.repo / 'training_normalization.json').read_text()) | |
| self.display = json.loads((self.repo / 'display_normalization.json').read_text()) | |
| self.class_names = json.loads((self.repo / 'classes.json').read_text())['names'] | |
| self.names = self.norm['score_names'] | |
| if self.names != [r['name'] for r in self.display['targets']]: | |
| raise ValueError('Display and training target orders differ') | |
| directory = self.repo / model_size | |
| self.model = LayeredMultiTaskWhisper(directory, len(self.class_names), initialize_pretrained=False) | |
| self.model.load_state_dict(load_file(directory / 'model.safetensors', device='cpu'), strict=True) | |
| self.model.to(device).eval() | |
| self.extractor = WhisperFeatureExtractor.from_pretrained(directory, local_files_only=True) | |
| def predict(self, audio, *, frame_threshold=.5, event_threshold=.5, max_events=32, include_frames=False): | |
| if not 0 < frame_threshold < 1 or not 0 < event_threshold < 1 or max_events < 1: | |
| raise ValueError('Thresholds must be between zero and one, max_events positive') | |
| wave, duration = read_audio(audio) | |
| features = self.extractor([wave], sampling_rate=16000, padding='longest', truncation=True, | |
| max_length=480000, return_attention_mask=True, return_tensors='np') | |
| mel = features['input_features'].astype(np.float32) | |
| mask = features['attention_mask'].astype(np.int64) | |
| if mel.shape[-1] % 2: | |
| mel = np.pad(mel, ((0, 0), (0, 0), (0, 1))) | |
| mask = np.pad(mask, ((0, 0), (0, 1))) | |
| from contextlib import nullcontext | |
| autocast = torch.autocast('cuda', dtype=torch.bfloat16) if self.device.startswith('cuda') else nullcontext() | |
| with torch.inference_mode(), autocast: | |
| out = self.model(torch.from_numpy(mel).to(self.device), torch.from_numpy(mask).to(self.device), | |
| predict_events=True, event_threshold=event_threshold, max_pred_events=max_events) | |
| standardized = out['scores'][0].float().cpu().numpy() | |
| raw = standardized * np.asarray(self.norm['score_std']) + np.asarray(self.norm['score_mean']) | |
| z = {name: ((float(value) - row['median']) / row['std'] | |
| if row['median'] is not None and row['std'] not in (None, 0) else None) | |
| for name, value, row in zip(self.names, raw, self.display['targets'])} | |
| count = min(int((mask.sum() + 1) // 2), out['frame'].shape[1]) | |
| frame = out['frame'][0, :count].float().sigmoid().cpu().numpy() | |
| edges = np.diff(np.pad((frame >= frame_threshold).astype(np.int8), (1, 1))) | |
| regions = [{'start_s': round(float(a) * .02, 6), 'end_s': round(min(float(b) * .02, duration), 6)} | |
| for a, b in zip(np.flatnonzero(edges == 1), np.flatnonzero(edges == -1)) | |
| if float(a) * .02 < duration] | |
| probabilities = out['event_class'][0].float().softmax(-1).cpu().numpy() | |
| starts = out['predicted_event_starts'][0].cpu().numpy() | |
| ends = out['predicted_event_ends'][0].cpu().numpy() | |
| valid = out['predicted_event_valid'][0].cpu().numpy() | |
| onset = out['onset'][0].float().sigmoid().cpu().numpy() | |
| proposals = [] | |
| for j in np.flatnonzero(valid): | |
| start, end = float(starts[j]) * .02, min(float(ends[j]) * .02, duration) | |
| if start >= end: | |
| continue | |
| top = np.argsort(probabilities[j])[-3:][::-1] | |
| proposals.append({'start_s': round(start, 6), 'end_s': round(end, 6), | |
| 'onset_probability': float(onset[int(starts[j])]), | |
| 'top3_classes': [{'label': self.class_names[int(k)], | |
| 'softmax_probability': float(probabilities[j, k])} for k in top]}) | |
| cps_z = float(out['cps'][0].float().cpu()) | |
| result = { | |
| 'model': self.model_size, 'checkpoint': 'best Gemini Flash 3.8 fine-tune; epoch 2', | |
| 'duration_s': duration, | |
| 'scores_raw': dict(zip(self.names, map(float, raw))), | |
| 'scores_training_z': dict(zip(self.names, map(float, standardized))), | |
| 'scores_z': z, | |
| 'characters_per_second': cps_z * self.norm['cps_std'] + self.norm['cps_mean'], | |
| 'characters_per_second_training_z': cps_z, | |
| 'estimated_burst_count': float(np.expm1(np.clip(raw[130], 0, 8))), | |
| 'orange_timbre_128d': out['timbre'][0].float().cpu().tolist(), | |
| 'orange_identity_250d': out['identity'][0].float().cpu().tolist(), | |
| 'binary_burst_regions': regions, 'burst_event_proposals': proposals, | |
| 'decoding': {'frame_threshold': frame_threshold, 'event_threshold': event_threshold, 'max_events': max_events}, | |
| 'notes': 'Encoder only: no transcript or captions. Raw regression scores are not clipped to ordinal scales. ' | |
| 'Display z-scores use the original train-only reference median/std; they are not probabilities. ' | |
| 'Class probabilities are uncalibrated. Blend is only meaningful when vocal bursts are present.', | |
| } | |
| if include_frames: | |
| result['burst_probability_per_20ms_frame'] = frame.tolist() | |
| return result | |
| def predict(audio, repo_dir, *, model_size='medium', device='cpu', threads=4, **options): | |
| """Convenience call; reuse Predictor for many clips to avoid reloading weights.""" | |
| return Predictor(repo_dir, model_size, device, threads).predict(audio, **options) | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('audio', type=Path) | |
| parser.add_argument('--repo-dir', type=Path, default=Path(__file__).resolve().parent) | |
| parser.add_argument('--model', choices=('base', 'medium', 'small'), default='medium', | |
| help='Release variant (small is a legacy alias for medium)') | |
| parser.add_argument('--device', default='cpu') | |
| parser.add_argument('--threads', type=int, default=4) | |
| parser.add_argument('--frame-threshold', type=float, default=.5) | |
| parser.add_argument('--event-threshold', type=float, default=.5) | |
| parser.add_argument('--max-events', type=int, default=32) | |
| parser.add_argument('--include-frame-probabilities', action='store_true') | |
| args = parser.parse_args() | |
| result = predict(args.audio, args.repo_dir, model_size=args.model, device=args.device, threads=args.threads, | |
| frame_threshold=args.frame_threshold, event_threshold=args.event_threshold, | |
| max_events=args.max_events, include_frames=args.include_frame_probabilities) | |
| print(json.dumps(result, indent=2, ensure_ascii=False, allow_nan=False)) | |
| if __name__ == '__main__': | |
| main() | |