Skip to content

utils

filter_by_class_count(entries, label_source, min_count)

Drop entries whose label_source value occurs <= min_count times.

A min_count of None (or <= 0) disables filtering.

Source code in mantra/datasets/utils.py
40
41
42
43
44
45
46
47
48
49
50
def filter_by_class_count(entries, label_source, min_count):
    """Drop entries whose ``label_source`` value occurs <= ``min_count`` times.

    A ``min_count`` of ``None`` (or <= 0) disables filtering.
    """
    if min_count is None or min_count <= 0:
        return entries, Counter()
    counts = Counter(e[label_source] for e in entries)
    kept_labels = {lbl for lbl, c in counts.items() if c > min_count}
    filtered = [e for e in entries if e[label_source] in kept_labels]
    return filtered, counts