Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
123 changes: 63 additions & 60 deletions src/main/python/tests/scuro/test_fusion_orders.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,77 +19,80 @@
#
# -------------------------------------------------------------

import os
import shutil
import unittest
import numpy as np

from systemds.scuro import Concatenation, RowMax, Hadamard
from systemds.scuro.modality.unimodal_modality import UnimodalModality
from systemds.scuro.representations.bert import Bert
from systemds.scuro.representations.mel_spectrogram import MelSpectrogram
from systemds.scuro.representations.average import Average
from tests.scuro.data_generator import ModalityRandomDataGenerator
from systemds.scuro.modality.type import ModalityType


class TestFusionOrders(unittest.TestCase):
"""
Tests the order properties of the fusion operators: commutativity, whether
the order of a pairwise chain matters, and whether a pairwise chain gives
the same result as the n-ary call.
"""

# (operator, chain_order_independent, chain_equals_nary)
#
# Commutativity is not in the table. Every Fusion operator has a
# "commutative" attribute. It is False on the base class and a subclass
# can override it. The test compares the measured result with that
# attribute, so an operator whose attribute does not match its
# implementation fails here.
#
# Combining a pair is never the same as combining all three. That is
# checked for every operator, so it is not in the table either.
FUSION_PROPERTIES = [
(Average, True, False),
(Concatenation, False, True),
(RowMax, True, True),
(Hadamard, True, True),
]

@classmethod
def setUpClass(cls):
cls.num_instances = 40
# These properties do not depend on the input shape, so the data can
# be small.
cls.num_instances = 4
cls.num_features = 8
cls.data_generator = ModalityRandomDataGenerator()
cls.r_1 = cls.data_generator.create1DModality(40, 100, ModalityType.AUDIO)
cls.r_2 = cls.data_generator.create1DModality(40, 100, ModalityType.TEXT)
cls.r_3 = cls.data_generator.create1DModality(40, 100, ModalityType.TEXT)

def test_fusion_order_avg(self):
r_1_r_2 = self.r_1.combine(self.r_2, Average())
r_2_r_1 = self.r_2.combine(self.r_1, Average())
r_1_r_2_r_3 = r_1_r_2.combine(self.r_3, Average())
r_2_r_1_r_3 = r_2_r_1.combine(self.r_3, Average())

r1_r2_r3 = self.r_1.combine([self.r_2, self.r_3], Average())

self.assertTrue(np.array_equal(r_1_r_2.data, r_2_r_1.data))
self.assertTrue(np.array_equal(r_1_r_2_r_3.data, r_2_r_1_r_3.data))
self.assertFalse(np.array_equal(r_1_r_2_r_3.data, r1_r2_r3.data))
self.assertFalse(np.array_equal(r_1_r_2.data, r1_r2_r3.data))

def test_fusion_order_concat(self):
r_1_r_2 = self.r_1.combine(self.r_2, Concatenation())
r_2_r_1 = self.r_2.combine(self.r_1, Concatenation())
r_1_r_2_r_3 = r_1_r_2.combine(self.r_3, Concatenation())
r_2_r_1_r_3 = r_2_r_1.combine(self.r_3, Concatenation())

r1_r2_r3 = self.r_1.combine([self.r_2, self.r_3], Concatenation())

self.assertFalse(np.array_equal(r_1_r_2.data, r_2_r_1.data))
self.assertFalse(np.array_equal(r_1_r_2_r_3.data, r_2_r_1_r_3.data))
self.assertFalse(np.array_equal(r_2_r_1.data, r1_r2_r3.data))
self.assertFalse(np.array_equal(r_1_r_2.data, r1_r2_r3.data))

def test_fusion_order_max(self):
r_1_r_2 = self.r_1.combine(self.r_2, RowMax())
r_2_r_1 = self.r_2.combine(self.r_1, RowMax())
r_1_r_2_r_3 = r_1_r_2.combine(self.r_3, RowMax())
r_2_r_1_r_3 = r_2_r_1.combine(self.r_3, RowMax())

r1_r2_r3 = self.r_1.combine([self.r_2, self.r_3], RowMax())

self.assertTrue(np.array_equal(r_1_r_2.data, r_2_r_1.data))
self.assertTrue(np.array_equal(r_1_r_2_r_3.data, r_2_r_1_r_3.data))
self.assertTrue(np.array_equal(r_1_r_2_r_3.data, r1_r2_r3.data))
self.assertFalse(np.array_equal(r_1_r_2.data, r1_r2_r3.data))

def test_fusion_order_hadamard(self):
r_1_r_2 = self.r_1.combine(self.r_2, Hadamard())
r_2_r_1 = self.r_2.combine(self.r_1, Hadamard())
r_1_r_2_r_3 = r_1_r_2.combine(self.r_3, Hadamard())
r_2_r_1_r_3 = r_2_r_1.combine(self.r_3, Hadamard())

r1_r2_r3 = self.r_1.combine([self.r_2, self.r_3], Hadamard())

self.assertTrue(np.array_equal(r_1_r_2.data, r_2_r_1.data))
self.assertTrue(np.array_equal(r_1_r_2_r_3.data, r_2_r_1_r_3.data))
self.assertTrue(np.array_equal(r_1_r_2_r_3.data, r1_r2_r3.data))
self.assertFalse(np.array_equal(r_1_r_2.data, r1_r2_r3.data))
def setUp(self):
self.r_1 = self.data_generator.create1DModality(
self.num_instances, self.num_features, ModalityType.AUDIO
)
self.r_2 = self.data_generator.create1DModality(
self.num_instances, self.num_features, ModalityType.TEXT
)
self.r_3 = self.data_generator.create1DModality(
self.num_instances, self.num_features, ModalityType.TEXT
)

@staticmethod
def _equal(left, right):
return np.array_equal(np.asarray(left.data), np.asarray(right.data))

def test_fusion_order_properties(self):
for (
fusion_operator,
chain_order_independent,
chain_equals_nary,
) in self.FUSION_PROPERTIES:
with self.subTest(fusion=fusion_operator.__name__):
r_1_r_2 = self.r_1.combine(self.r_2, fusion_operator())
r_2_r_1 = self.r_2.combine(self.r_1, fusion_operator())
r_1_r_2_r_3 = r_1_r_2.combine(self.r_3, fusion_operator())
r_2_r_1_r_3 = r_2_r_1.combine(self.r_3, fusion_operator())
r1_r2_r3 = self.r_1.combine([self.r_2, self.r_3], fusion_operator())

self.assertEqual(
self._equal(r_1_r_2, r_2_r_1), fusion_operator().commutative
)
self.assertEqual(
self._equal(r_1_r_2_r_3, r_2_r_1_r_3), chain_order_independent
)
self.assertEqual(self._equal(r_1_r_2_r_3, r1_r2_r3), chain_equals_nary)
self.assertFalse(self._equal(r_1_r_2, r1_r2_r3))
Loading