Skip to content
Merged
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
120 changes: 120 additions & 0 deletions tinyml-modelmaker/tests/test_argv_builder_transform_bugs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,120 @@
"""Regression tests for three argv-builder bugs in vision/audio ai_modules:

1. image_base.py's train argv rendered data_proc_transforms as a stringified
list while test argv passed the raw list -- prepare_transforms() (in
tinyml-tinyverse) only combines it into args.transforms when
isinstance(args.data_proc_transforms, list) is true, so the stringified
form silently skipped it during training while testing still applied it.

2. audio_base.py's train and test argv builders each declared
--data-proc-transforms/--feat-ext-transform twice (once raw, once
stringified); argparse's last-occurrence-wins semantics meant the earlier,
correct raw declaration was always shadowed by the later, stringified one.

3. Fixing (1) by making data_proc_transforms a raw list exposed a second bug
an independent peer review caught: prepare_transforms() does
`args.data_proc_transforms + args.feat_ext_transform` whenever
data_proc_transforms is a list. image_base.py's train argv builder still
stringified --feat-ext-transform (and --augmentation-transform), so this
became `list + str`, raising TypeError on every image-classification
training run. test_image_train_argv_feat_ext_transform_survives_
prepare_transforms below exercises the REAL prepare_transforms() against
the argv builder's actual output -- the shape-only MagicMock tests below
couldn't catch this since they never fed argv through it.

Most of these are pure unit tests against the argv-building methods in
isolation (constructed via unittest.mock.MagicMock for self, rather than a
full ModelRunner instance) -- these methods only read attributes off
self.params and don't need real training infrastructure.
"""
from argparse import Namespace
from unittest.mock import MagicMock

from tinyml_modelmaker.ai_modules.vision.training.tinyml_tinyverse.image_base import (
BaseImageModelTraining,
)
from tinyml_modelmaker.ai_modules.audio.training.tinyml_tinyverse.audio_base import (
BaseAudioModelTraining,
)
from tinyml_tinyverse.references.common.train_base import prepare_transforms


def _argv_value_after(argv, flag):
return argv[argv.index(flag) + 1]


def _count_occurrences(argv, flag):
return argv.count(flag)


def test_image_train_argv_passes_data_proc_transforms_as_a_raw_list():
fake_self = MagicMock()
fake_self.params.data_processing_feature_extraction.data_proc_transforms = ["BINARIZE"]

argv = BaseImageModelTraining._build_common_train_argv(fake_self, device="cpu", distributed=0)

value = _argv_value_after(argv, "--data-proc-transforms")
assert value == ["BINARIZE"]
assert isinstance(value, list)


def test_image_train_argv_feat_ext_transform_survives_prepare_transforms():
"""Drives the REAL prepare_transforms() (tinyml-tinyverse) against the
actual argv the train-argv builder produces -- data_proc_transforms and
feat_ext_transform must both come out as lists, or the `+` inside
prepare_transforms raises TypeError."""
fake_self = MagicMock()
fake_self.params.data_processing_feature_extraction.data_proc_transforms = ["BINARIZE"]
fake_self.params.data_processing_feature_extraction.feat_ext_transform = ["MFCC"]

argv = BaseImageModelTraining._build_common_train_argv(fake_self, device="cpu", distributed=0)

args = Namespace(
data_proc_transforms=_argv_value_after(argv, "--data-proc-transforms"),
feat_ext_transform=_argv_value_after(argv, "--feat-ext-transform"),
)

prepare_transforms(args)

assert args.transforms == ["BINARIZE", "MFCC"]


def test_image_train_and_test_argv_agree_on_data_proc_transforms_form():
fake_self = MagicMock()
fake_self.params.data_processing_feature_extraction.data_proc_transforms = ["BINARIZE", "RESIZE"]

train_argv = BaseImageModelTraining._build_common_train_argv(fake_self, device="cpu", distributed=0)
test_argv = BaseImageModelTraining._build_common_test_argv(
fake_self, device="cpu", data_path="/tmp/data", model_path="/tmp/model.onnx", output_dir="/tmp/out"
)

assert _argv_value_after(train_argv, "--data-proc-transforms") == \
_argv_value_after(test_argv, "--data-proc-transforms")


def test_audio_train_argv_declares_data_proc_transforms_exactly_once():
fake_self = MagicMock()
fake_self.params.data_processing_feature_extraction.data_proc_transforms = ["NORMALIZE"]
fake_self.params.data_processing_feature_extraction.feat_ext_transform = ["MFCC"]

argv = BaseAudioModelTraining._build_common_train_argv(fake_self, device="cpu", distributed=0)

assert _count_occurrences(argv, "--data-proc-transforms") == 1
assert _count_occurrences(argv, "--feat-ext-transform") == 1
assert _argv_value_after(argv, "--data-proc-transforms") == ["NORMALIZE"]
assert _argv_value_after(argv, "--feat-ext-transform") == ["MFCC"]


def test_audio_test_argv_declares_data_proc_transforms_exactly_once():
fake_self = MagicMock()
fake_self.params.data_processing_feature_extraction.data_proc_transforms = ["NORMALIZE"]
fake_self.params.data_processing_feature_extraction.feat_ext_transform = ["MFCC"]

argv = BaseAudioModelTraining._build_common_test_argv(
fake_self, device="cpu", data_path="/tmp/data", model_path="/tmp/model.onnx", output_dir="/tmp/out"
)

assert _count_occurrences(argv, "--data-proc-transforms") == 1
assert _count_occurrences(argv, "--feat-ext-transform") == 1
assert _argv_value_after(argv, "--data-proc-transforms") == ["NORMALIZE"]
assert _argv_value_after(argv, "--feat-ext-transform") == ["MFCC"]
Original file line number Diff line number Diff line change
Expand Up @@ -341,9 +341,6 @@ def _build_common_train_argv(self, device, distributed):
'--normalize-audio', f'{self.params.data_processing_feature_extraction.normalize_audio}',
'--mono', f'{self.params.data_processing_feature_extraction.mono}',

'--data-proc-transforms', f'{self.params.data_processing_feature_extraction.data_proc_transforms}',
'--feat-ext-transform', f'{self.params.data_processing_feature_extraction.feat_ext_transform}',

'--output-int', f'{self.params.training.output_int}',
'--variables', f'{self.params.data_processing_feature_extraction.variables}',
'--lis', f'{self.params.training.log_file_path}',
Expand Down Expand Up @@ -392,8 +389,6 @@ def _build_common_test_argv(self, device, data_path, model_path, output_dir):
'--normalize-audio', f'{self.params.data_processing_feature_extraction.normalize_audio}',
'--mono', f'{self.params.data_processing_feature_extraction.mono}',

'--data-proc-transforms', f'{self.params.data_processing_feature_extraction.data_proc_transforms}',
'--feat-ext-transform', f'{self.params.data_processing_feature_extraction.feat_ext_transform}',
'--nn-for-feature-extraction', f'{self.params.data_processing_feature_extraction.nn_for_feature_extraction}',
'--output-int', f'{self.params.training.output_int}',

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -339,9 +339,18 @@ def _build_common_train_argv(self, device, distributed):
'--generic-model', f'{self.params.common.generic_model}',
'--sampling-rate', f'{self.params.data_processing_feature_extraction.sampling_rate}',
# Transform
'--data-proc-transforms', f'{self.params.data_processing_feature_extraction.data_proc_transforms}',
'--feat-ext-transform', f'{self.params.data_processing_feature_extraction.feat_ext_transform}',
'--augmentation-transform', f'{self.params.data_processing_feature_extraction.augmentation_transform}',
# Pass raw lists (not stringified) -- matches the test argv builder
# below and timeseries_base.py's reference implementation.
# prepare_transforms() (train_base.py) does
# `args.data_proc_transforms + args.feat_ext_transform` whenever
# isinstance(args.data_proc_transforms, list) is true. Both operands
# must therefore be actual lists at that point, not just
# data_proc_transforms: a stringified feat_ext_transform (e.g. "[]")
# makes this `list + str`, raising TypeError on every training run
# once data_proc_transforms itself was made a raw list.
'--data-proc-transforms', self.params.data_processing_feature_extraction.data_proc_transforms,
'--feat-ext-transform', self.params.data_processing_feature_extraction.feat_ext_transform,
'--augmentation-transform', self.params.data_processing_feature_extraction.augmentation_transform,
'--feat-ext-store-dir', f'{self.params.data_processing_feature_extraction.feat_ext_store_dir}',
'--dont-train-just-feat-ext', f'{self.params.data_processing_feature_extraction.dont_train_just_feat_ext}',
'--store-feat-ext-data', f'{self.params.data_processing_feature_extraction.store_feat_ext_data}',
Expand Down
Loading