ChristophSchuhmann's picture
Brand release as Humaneness Ears Base and Medium and update inference
c1cc83c verified
Raw History Blame Contribute Delete
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()