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
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
setup_training_environment,
prepare_transforms,
create_data_loaders,
shutdown_data_loaders,
# Model creation and setup
create_model,
log_model_summary,
Expand Down
86 changes: 86 additions & 0 deletions tinyml-tinyverse/tinyml_tinyverse/references/common/train_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -808,3 +808,89 @@ def create_data_loaders(dataset, dataset_test, train_sampler, test_sampler, args
dataset_test, batch_size=args.batch_size, sampler=test_sampler, num_workers=args.workers,
pin_memory=True if gpu > 0 else False, collate_fn=utils.collate_fn)
return data_loader, data_loader_test


def _unregister_tracked_semaphores(*objects):
"""Unregister multiprocessing semaphores from Python's resource_tracker.

On Python <=3.11, ``_multiprocessing.SemLock``'s C dealloc calls
``sem_close()`` but never ``resource_tracker.unregister()``. The
resource_tracker's ``atexit`` handler therefore reports every
semaphore ever created as "leaked". (Fixed in Python 3.12+ where
SemLock.__del__ calls unregister.)

This function walks known multiprocessing container attributes
(Queue._rlock, Queue._wlock, Queue._sem, Event._cond, Event._flag,
etc.) to find underlying SemLock objects and unregisters their named
semaphores so the resource_tracker stays quiet.
"""
try:
from multiprocessing.resource_tracker import unregister
except ImportError:
return

seen = set()

def _scan(obj):
obj_id = id(obj)
if obj_id in seen:
return
seen.add(obj_id)
# Leaf: a Lock / Semaphore / BoundedSemaphore wrapping a SemLock
semlock = getattr(obj, '_semlock', None)
if semlock is not None:
name = getattr(semlock, 'name', None)
if name:
try:
unregister(name, "semaphore")
except Exception:
pass
return
# Recurse into known container attributes
for attr in ('_rlock', '_wlock', '_sem', '_lock', '_cond', '_flag'):
child = getattr(obj, attr, None)
if child is not None:
_scan(child)

for obj in objects:
if obj is None:
continue
if isinstance(obj, (list, tuple)):
for item in obj:
_scan(item)
else:
_scan(obj)


def shutdown_data_loaders(*loaders):
"""Explicitly shut down DataLoader worker processes to avoid leaked semaphore warnings.

Must be called before exit when DataLoaders use num_workers > 0 (especially
on macOS where the 'spawn' start method tracks semaphores via resource_tracker).
Works for both persistent_workers=True and False.

After joining workers, we explicitly unregister all POSIX named semaphores
owned by the iterator's multiprocessing Queues and Events from the
resource_tracker. On Python <=3.11, this unregister never happens
automatically (the C SemLock dealloc only calls sem_close, not
resource_tracker.unregister), so without this step the resource_tracker
warns about "leaked semaphore objects" at shutdown.
"""
import gc
for loader in loaders:
if not (hasattr(loader, '_iterator') and loader._iterator is not None):
continue
it = loader._iterator
try:
it._shutdown_workers()
except Exception:
pass
# Unregister all POSIX named semaphores from the resource_tracker
# so it does not report them as leaked at exit.
_unregister_tracked_semaphores(
getattr(it, '_index_queues', None),
getattr(it, '_data_queue', None),
getattr(it, '_workers_done_event', None),
)
loader._iterator = None
gc.collect()
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
from argparse import ArgumentParser
from logging import getLogger

import numpy as np
import onnxruntime as ort
import pandas as pd
import torch
Expand All @@ -47,6 +48,7 @@
from tinyml_tinyverse.common.utils import misc_utils, utils, mdcl_utils
from tinyml_tinyverse.common.utils.mdcl_utils import Logger
from tinyml_tinyverse.common.utils.utils import get_confusion_matrix
from ..common.train_base import shutdown_data_loaders

# Import common functions from base module
from ..common.test_onnx_base import (
Expand Down Expand Up @@ -131,73 +133,76 @@ def main(gpu, args):
data_loader = torch.utils.data.DataLoader(
dataset, batch_size=args.batch_size, sampler=train_sampler,
num_workers=args.workers, pin_memory=True, collate_fn=utils.collate_fn)

logger.info(f"Loading ONNX model: {args.model_path}")
ort_sess, input_name, output_name = load_onnx_model(args.model_path, args.generic_model)

predicted = torch.tensor([]).to(device, non_blocking=True)
ground_truth = torch.tensor([]).to(device, non_blocking=True)
for batched_raw_data, batched_data, batched_target in data_loader:
batched_raw_data = batched_raw_data.to(device, non_blocking=True).long()
batched_data = batched_data.to(device, non_blocking=True).float()
batched_target = batched_target.to(device, non_blocking=True).long()
if transform:
batched_data = transform(batched_data)
if args.nn_for_feature_extraction:
for data in batched_raw_data:
predicted = torch.cat((predicted, torch.tensor(
ort_sess.run([output_name], {input_name: data.unsqueeze(0).cpu().numpy().astype(np.float32)})[0]
).to(device)))
else:
for data in batched_data:
predicted = torch.cat((predicted, torch.tensor(
ort_sess.run([output_name], {input_name: data.unsqueeze(0).cpu().numpy()})[0]
).to(device)))
ground_truth = torch.cat((ground_truth, batched_target))

try:
mdcl_utils.create_dir(os.path.join(args.output_dir, 'post_training_analysis'))
logger.info("Plotting OvR Multiclass ROC score")
utils.plot_multiclass_roc(ground_truth, predicted, os.path.join(args.output_dir, 'post_training_analysis'),
label_map=dataset.inverse_label_map, phase='test')
logger.info("Plotting Class difference scores")
utils.plot_pairwise_differenced_class_scores(ground_truth, predicted,
os.path.join(args.output_dir, 'post_training_analysis'),
label_map=dataset.inverse_label_map, phase='test')
except Exception as e:
logger.warning(f"Post Training Analysis plots will not be generated because: {e}")

metric = torcheval.metrics.MulticlassAccuracy()
# predicted = torch.argmax(predicted, dim=1)
metric.update(predicted, ground_truth)
logger = getLogger("root.main.test_data")
logger.info(f"Test Data Evaluation Accuracy: {metric.compute() * 100:.2f}%")
try:
logger.info(
f"Test Data Evaluation AUC ROC Score: {utils.get_au_roc(predicted, ground_truth, num_classes):.3f}")
except ValueError as e:
logger.warning("Not able to compute AUC ROC. Error: " + str(e))
if len(torch.unique(ground_truth)) == 1:
logger.warning("Confusion Matrix can not be printed because only items of 1 class was present in test data")
else:
try:
confusion_matrix = get_confusion_matrix(predicted, ground_truth.type(torch.int64),
num_classes).cpu().numpy()
logger.info('Confusion Matrix:\n {}'.format(tabulate(pd.DataFrame(
confusion_matrix, columns=[f"Predicted as: {x}" for x in dataset.inverse_label_map.values()],
index=[f"Ground Truth: {x}" for x in dataset.inverse_label_map.values()]), headers="keys", tablefmt='grid')))
except ValueError as e:
logger.warning("Not able to compute Confusion Matrix. Error: " + str(e))

logger.info(f"Loading ONNX model: {args.model_path}")
ort_sess, input_name, output_name = load_onnx_model(args.model_path, args.generic_model)

predicted = torch.tensor([]).to(device, non_blocking=True)
ground_truth = torch.tensor([]).to(device, non_blocking=True)
for batched_raw_data, batched_data, batched_target in data_loader:
batched_raw_data = batched_raw_data.to(device, non_blocking=True).long()
batched_data = batched_data.to(device, non_blocking=True).float()
batched_target = batched_target.to(device, non_blocking=True).long()
if transform:
batched_data = transform(batched_data)
if args.nn_for_feature_extraction:
for data in batched_raw_data:
predicted = torch.cat((predicted, torch.tensor(
ort_sess.run([output_name], {input_name: data.unsqueeze(0).cpu().numpy().astype(np.float32)})[0]
).to(device)))
else:
for data in batched_data:
predicted = torch.cat((predicted, torch.tensor(
ort_sess.run([output_name], {input_name: data.unsqueeze(0).cpu().numpy()})[0]
).to(device)))
ground_truth = torch.cat((ground_truth, batched_target))

try:
Logger(log_file=args.file_level_classification_log, DEBUG=args.DEBUG,
name="root.utils.print_file_level_classification_summary", append_log=True, console_log=False)
getLogger("root.utils.print_file_level_classification_summary").propagate = False
utils.print_file_level_classification_summary(dataset_test, predicted, ground_truth, "TestData")
logger.info(f"Generated File-level classification summary of test data in: {args.file_level_classification_log}")
mdcl_utils.create_dir(os.path.join(args.output_dir, 'post_training_analysis'))
logger.info("Plotting OvR Multiclass ROC score")
utils.plot_multiclass_roc(ground_truth, predicted, os.path.join(args.output_dir, 'post_training_analysis'),
label_map=dataset.inverse_label_map, phase='test')
logger.info("Plotting Class difference scores")
utils.plot_pairwise_differenced_class_scores(ground_truth, predicted,
os.path.join(args.output_dir, 'post_training_analysis'),
label_map=dataset.inverse_label_map, phase='test')
except Exception as e:
logger.error(f"Failed to generate file-level classification summary: {str(e)}")
logger.warning(f"Post Training Analysis plots will not be generated because: {e}")

metric = torcheval.metrics.MulticlassAccuracy()
# predicted = torch.argmax(predicted, dim=1)
metric.update(predicted, ground_truth)
logger = getLogger("root.main.test_data")
logger.info(f"Test Data Evaluation Accuracy: {metric.compute() * 100:.2f}%")
try:
logger.info(
f"Test Data Evaluation AUC ROC Score: {utils.get_au_roc(predicted, ground_truth, num_classes):.3f}")
except ValueError as e:
logger.warning("Not able to compute AUC ROC. Error: " + str(e))
if len(torch.unique(ground_truth)) == 1:
logger.warning("Confusion Matrix can not be printed because only items of 1 class was present in test data")
else:
try:
confusion_matrix = get_confusion_matrix(predicted, ground_truth.type(torch.int64),
num_classes).cpu().numpy()
logger.info('Confusion Matrix:\n {}'.format(tabulate(pd.DataFrame(
confusion_matrix, columns=[f"Predicted as: {x}" for x in dataset.inverse_label_map.values()],
index=[f"Ground Truth: {x}" for x in dataset.inverse_label_map.values()]), headers="keys", tablefmt='grid')))
except ValueError as e:
logger.warning("Not able to compute Confusion Matrix. Error: " + str(e))

try:
Logger(log_file=args.file_level_classification_log, DEBUG=args.DEBUG,
name="root.utils.print_file_level_classification_summary", append_log=True, console_log=False)
getLogger("root.utils.print_file_level_classification_summary").propagate = False
utils.print_file_level_classification_summary(dataset_test, predicted, ground_truth, "TestData")
logger.info(f"Generated File-level classification summary of test data in: {args.file_level_classification_log}")
except Exception as e:
logger.error(f"Failed to generate file-level classification summary: {str(e)}")

finally:
shutdown_data_loaders(data_loader)
return

def run(args):
Expand Down
Loading