Skip to content

Commit 0f14532

Browse files
authored
Revert "Balance family training data (#31)"
This reverts commit a060959.
1 parent a060959 commit 0f14532

3 files changed

Lines changed: 2 additions & 130 deletions

File tree

llm_fingerprinter/cli.py

Lines changed: 2 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,6 @@
2323
from llm_fingerprinter.fingerprinter import LLMFingerprinter
2424
from llm_fingerprinter.fingerprint_store import FingerprintStore
2525
from llm_fingerprinter.template_classifier import TemplateClassifier
26-
from llm_fingerprinter.training_data import balance_grouped_samples
2726
from llm_fingerprinter.plotting import FingerprintPlotError, plot_fingerprint_projection
2827
from llm_fingerprinter.contracts.llm import LLMRequest, Message
2928
from llm_fingerprinter.providers import create_provider
@@ -575,13 +574,8 @@ def simulate(ctx, backend, endpoint, api_key, request_file, model, family, num_s
575574
help='Run k-fold cross-validation before training')
576575
@click.option('--cv-folds', default=5, type=int,
577576
help='Number of cross-validation folds (default: 5)')
578-
@click.option('--balance/--no-balance', default=True, show_default=True,
579-
help='Downsample each family to the same number of fingerprints.')
580-
@click.option('--balance-seed', default=42, type=int, show_default=True,
581-
help='Random seed used for reproducible family balancing.')
582577
@click.pass_context
583-
def train(ctx, augment, use_pca, pca_components, cross_validate, cv_folds,
584-
balance, balance_seed):
578+
def train(ctx, augment, use_pca, pca_components, cross_validate, cv_folds):
585579
"""Train classifier from saved fingerprints.
586580
587581
\b
@@ -612,24 +606,6 @@ def train(ctx, augment, use_pca, pca_components, cross_validate, cv_folds,
612606
click.echo(" Run 'simulate' first")
613607
sys.exit(1)
614608

615-
if balance:
616-
original_counts = {
617-
family: len(vectors)
618-
for family, vectors in training_data.items()
619-
if vectors
620-
}
621-
training_data, target_count = balance_grouped_samples(
622-
training_data, seed=balance_seed
623-
)
624-
click.echo(
625-
f"\nBalanced each family to {target_count} samples "
626-
f"(seed={balance_seed}):"
627-
)
628-
for family, vectors in sorted(training_data.items()):
629-
click.echo(
630-
f" {family}: {original_counts[family]} -> {len(vectors)}"
631-
)
632-
633609
click.echo("\n📊 Training data:")
634610
total = 0
635611
for fam, vecs in sorted(training_data.items()):
@@ -1045,12 +1021,8 @@ def fingerprint(ctx, backend, endpoint, api_key, request_file, model, repeats, o
10451021
@cli.command('build-templates')
10461022
@click.option('--ood-ratio', default=0.80, type=float, show_default=True,
10471023
help='OOD ratio threshold (lower = stricter OOD detection).')
1048-
@click.option('--balance/--no-balance', default=True, show_default=True,
1049-
help='Downsample each family to the same number of fingerprints.')
1050-
@click.option('--balance-seed', default=42, type=int, show_default=True,
1051-
help='Random seed used for reproducible family balancing.')
10521024
@click.pass_context
1053-
def build_templates(ctx, ood_ratio, balance, balance_seed):
1025+
def build_templates(ctx, ood_ratio):
10541026
"""Build open-set template classifier from training fingerprints.
10551027
10561028
Templates let you classify new families without retraining the ensemble —
@@ -1078,24 +1050,6 @@ def build_templates(ctx, ood_ratio, balance, balance_seed):
10781050
click.echo(" Run 'simulate' first")
10791051
sys.exit(1)
10801052

1081-
if balance:
1082-
original_counts = {
1083-
family: len(vectors)
1084-
for family, vectors in training_data.items()
1085-
if vectors
1086-
}
1087-
training_data, target_count = balance_grouped_samples(
1088-
training_data, seed=balance_seed
1089-
)
1090-
click.echo(
1091-
f"\nBalanced each family to {target_count} samples "
1092-
f"(seed={balance_seed}):"
1093-
)
1094-
for family, vectors in sorted(training_data.items()):
1095-
click.echo(
1096-
f" {family:12s} {original_counts[family]} -> {len(vectors)}"
1097-
)
1098-
10991053
click.echo("\n📊 Training data:")
11001054
for fam, vecs in sorted(training_data.items()):
11011055
click.echo(f" {fam:12s} {len(vecs)} samples")

llm_fingerprinter/training_data.py

Lines changed: 0 additions & 41 deletions
This file was deleted.

tests/training_data/test_training_data.py

Lines changed: 0 additions & 41 deletions
This file was deleted.

0 commit comments

Comments
 (0)