"""Each torch-free branch must reproduce its training-time pipeline.""" import os import sys import numpy as np import pytest from reachy_mini_multimodal_emotion.engine import load_config from reachy_mini_multimodal_emotion.engine.face import FaceBranch, crop_face, preprocess_face from reachy_mini_multimodal_emotion.engine.speech import LogMel # ------------------------------------------------------------------ speech @pytest.mark.parametrize("seconds", [0.5, 2.0, 3.3]) def test_numpy_log_mel_matches_torchaudio(seconds): torch = pytest.importorskip("torch") T = pytest.importorskip("torchaudio.transforms") rng = np.random.default_rng(0) t = np.arange(int(seconds * 16000)) / 16000 wave = (0.1 * np.sin(2 * np.pi * 180 * t) + 0.02 * rng.standard_normal(len(t))).astype(np.float32) mel = T.MelSpectrogram(sample_rate=16000, n_fft=512, win_length=400, hop_length=160, n_mels=80, f_min=20.0, f_max=8000.0, center=True, power=2.0)(torch.from_numpy(wave)) log_mel = torch.log(torch.clamp(mel, min=1e-5)) ref = ((log_mel - log_mel.mean(-1, keepdim=True)) / (log_mel.std(-1, keepdim=True) + 1e-5)).numpy() ours = LogMel()(wave) assert ours.shape == ref.shape and np.abs(ours - ref).max() < 1e-3 # ------------------------------------------------------------------ face def test_face_preprocess_matches_torchvision_transform(): pytest.importorskip("torchvision") import cv2 from PIL import Image from torchvision import transforms rng = np.random.default_rng(1) face = rng.integers(0, 256, (97, 83, 3), dtype=np.uint8) gray = cv2.cvtColor(face, cv2.COLOR_BGR2GRAY) tf = transforms.Compose([transforms.Resize((112, 112)), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))]) ref = tf(Image.fromarray(np.repeat(gray[..., None], 3, axis=2), mode="RGB")).unsqueeze(0).numpy() np.testing.assert_allclose(preprocess_face(face), ref, atol=1e-5) def test_crop_face_pads_and_clamps(): frame = np.zeros((100, 200, 3), np.uint8) assert crop_face(frame, (10, 10, 50, 50), 0.10).shape == (60, 60, 3) assert crop_face(frame, (0, 0, 50, 50), 0.10).shape == (55, 55, 3) NEW_EMOTION = os.environ.get("NEW_EMOTION_REPO", "D:/newEmotion") def test_face_onnx_matches_training_checkpoint(model_root): torch = pytest.importorskip("torch") ckpt = os.path.join(NEW_EMOTION, "runs", "student_direct", "best.pt") if not os.path.exists(ckpt): pytest.skip("newEmotion checkpoint not available") sys.path.insert(0, NEW_EMOTION) from emotion_model.models import create_model payload = torch.load(ckpt, map_location="cpu", weights_only=False) model = create_model("mobilenet_v3_large", num_classes=7, pretrained=False) model.load_state_dict(payload["model_state"]) model.eval() cfg = load_config(model_root) branch = FaceBranch(model_root / cfg["files"]["face"], model_root / cfg["files"]["face_detector"]) face = np.random.default_rng(2).integers(0, 256, (120, 100, 3), dtype=np.uint8) with torch.no_grad(): ref = model(torch.from_numpy(preprocess_face(face))).numpy()[0] np.testing.assert_allclose(branch.logits(face), ref, atol=1e-3) # ------------------------------------------------------------------ text def test_fast_tokenizer_matches_transformers(model_root): transformers = pytest.importorskip("transformers") from tokenizers import Tokenizer src = os.environ.get("TER_MODEL", "D:/TER/TextEmotionDetection-model-2026-09-23/models/best_model") if not os.path.exists(src): pytest.skip("original text model not available") ref = transformers.AutoTokenizer.from_pretrained(src) ours = Tokenizer.from_file(str(model_root / "text" / "tokenizer.json")) ours.enable_truncation(128) for text in ["我今天終於完成專案了,超開心!", "他說:「我不是不生氣」", "OK啦 123 abc", "很" * 300]: enc = ours.encode(text) r = ref(text, truncation=True, max_length=128) assert enc.ids == r["input_ids"] and enc.type_ids == r["token_type_ids"] def test_text_branch_statuses(model_root): pytest.importorskip("chinese_converter") from reachy_mini_multimodal_emotion.engine.text import TTL_S, TextBranch cfg = load_config(model_root) tb = TextBranch(model_root / cfg["files"]["text"], model_root / cfg["files"]["text_tokenizer"]) assert tb.classify("", "zh", 0.0).status == "no_transcript" assert tb.classify("I am so happy today", "en", 0.0).status == "unsupported_language" r = tb.classify("我今天终于完成项目了,超开心!", "zh", 10.0) # Simplified input is converted assert r.present and tb.transcript.startswith("我今天終於") assert int(r.probs.argmax()) == 1 # happy assert tb.current(10.0 + TTL_S - 0.1).present assert tb.current(10.0 + TTL_S + 0.1).status == "stale"