"""Label maps and late fusion behaviour, including missing modalities.""" import numpy as np from reachy_mini_multimodal_emotion.engine.fusion import OK, BranchResult, fuse from reachy_mini_multimodal_emotion.engine.labels import BRANCH_LABELS, EMOTIONS, softmax, to_canonical def onehot_logits(branch, emotion, high=5.0): z = np.zeros(7) z[BRANCH_LABELS[branch].index(emotion)] = high return z def test_each_branch_reorders_into_canonical_labels(): for branch in BRANCH_LABELS: for emotion in EMOTIONS: assert EMOTIONS[int(to_canonical(branch, onehot_logits(branch, emotion)).argmax())] == emotion def test_face_reorder_matches_ser_handoff_index_list(): # SER_models/README.md: reorder face logits into speech order with [6, 3, 4, 0, 5, 2, 1] z = np.arange(7.0) assert to_canonical("face", z).tolist() == z[[6, 3, 4, 0, 5, 2, 1]].tolist() def probs_for(emotion, p): x = np.full(7, (1 - p) / 6) x[EMOTIONS.index(emotion)] = p return x W = {"speech": 0.35, "face": 0.65, "text": 0.35} def test_missing_branches_are_dropped_not_counted_as_neutral(): only_speech = fuse({"speech": BranchResult(probs_for("angry", 0.6), OK), "face": BranchResult(status="no_face"), "text": BranchResult(status="unsupported_language")}, W) assert only_speech.used == {"speech": 1.0} np.testing.assert_allclose(only_speech.probs, probs_for("angry", 0.6)) def test_nothing_present_gives_no_prediction(): fused = fuse({"speech": BranchResult(status="silence"), "face": BranchResult(status="no_face")}, W) assert fused.probs is None and fused.top is None and fused.used == {} def test_log_linear_pooling_weights_and_normalisation(): a, b = probs_for("happy", 0.7), probs_for("sad", 0.7) fused = fuse({"speech": BranchResult(a, OK), "face": BranchResult(b, OK)}, W) expected = np.exp(0.35 * np.log(a) + 0.65 * np.log(b)) np.testing.assert_allclose(fused.probs, expected / expected.sum()) assert fused.top == "sad" and abs(sum(fused.used.values()) - 1) < 1e-12 def test_agreeing_branches_raise_confidence(): p = probs_for("surprise", 0.5) fused = fuse({b: BranchResult(p, OK) for b in W}, W) assert fused.top == "surprise" and np.isclose(fused.confidence, 0.5) # identical inputs: unchanged def test_temperature_flattens(): z = np.array([3.0, 0, 0, 0, 0, 0, 0]) assert softmax(z, 6.46).max() < softmax(z, 1.0).max()