fix(backend): group augmented images with source photo in train/val split; document classifier retrain effort (task 2.5)
train_classifier.py's split_dataset() previously shuffled and split individual image files, letting an augmented copy (photo_aug_2.jpeg) land in validation while its near-duplicate source stayed in training - inflating val accuracy with memorization rather than measuring real generalization. Now groups by source photo (stripping _aug_N) before shuffling and splitting 80/20. Also records the in-progress effort to retrain the product classifier against the full 81-class/2,493-photo foto-kemasan-v2 dataset (up from the 16 classes/118 photos the deployed model was actually trained on) - see plans/next-enhancements.md task 2.5 and the accompanying iteration-log entry for the real, currently-observed numbers (DINOv2 index rebuilt: 2493/2493 images; classifier training: in progress, ~32s/epoch observed). Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01Xsxk4ZkDQVVaLUcixDcqb5
This commit is contained in:
1 parent
3a17c28758
commit
f4ec541369
3 files changed
+222
-14
No files matched your search
@@ -73,16 +73,28 @@ def latest_classifier_weights(models_dir: Path = DEFAULT_MODELS_DIR) -> Path:
|
||||
return max(candidates, key=sort_key)
|
||||
|
||||
VALID_IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
|
||||
AUG_SUFFIX_RE = re.compile(r"_aug_\d+$")
|
||||
|
||||
|
||||
def is_image_file(path: Path) -> bool:
|
||||
return path.is_file() and path.suffix.lower() in VALID_IMAGE_EXTENSIONS
|
||||
|
||||
|
||||
def _source_group_key(filename_stem: str) -> str:
|
||||
"""Strip an `_aug_<n>` suffix so an augmented image groups with its source photo."""
|
||||
return AUG_SUFFIX_RE.sub("", filename_stem)
|
||||
|
||||
|
||||
def split_dataset(src_dir: Path, dest_dir: Path, split_ratio: float = 0.8, seed: int = 42):
|
||||
"""
|
||||
Split class folders from src_dir into train/val folders in dest_dir.
|
||||
Ensures every class with 2+ images keeps at least one image in validation.
|
||||
|
||||
Splits by *source photo group*, not by individual file: an augmented image
|
||||
(`photo1_aug_2.jpeg`) always stays in the same split as its source
|
||||
(`photo1.jpeg`). Splitting file-by-file would let near-duplicate images
|
||||
land on opposite sides of train/val, inflating val accuracy with
|
||||
memorization instead of measuring generalization.
|
||||
"""
|
||||
random.seed(seed)
|
||||
|
||||
@@ -111,29 +123,41 @@ def split_dataset(src_dir: Path, dest_dir: Path, split_ratio: float = 0.8, seed:
|
||||
[f for f in c_dir.iterdir() if is_image_file(f)],
|
||||
key=lambda p: p.name,
|
||||
)
|
||||
random.shuffle(images)
|
||||
|
||||
num_images = len(images)
|
||||
if num_images == 0:
|
||||
print(f"Warning: Class '{class_name}' has 0 images. Skipping.")
|
||||
continue
|
||||
|
||||
# Group by source photo (stripping any `_aug_N` suffix) so an
|
||||
# augmented image and the photo it came from always land on the same
|
||||
# side of the split.
|
||||
groups: dict[str, list[Path]] = {}
|
||||
for img in images:
|
||||
groups.setdefault(_source_group_key(img.stem), []).append(img)
|
||||
group_keys = sorted(groups.keys())
|
||||
random.shuffle(group_keys)
|
||||
|
||||
class_train_dir = train_dir / class_name
|
||||
class_val_dir = val_dir / class_name
|
||||
class_train_dir.mkdir(parents=True, exist_ok=True)
|
||||
class_val_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if num_images == 1:
|
||||
train_images = images
|
||||
val_images = images
|
||||
elif num_images == 2:
|
||||
train_images = [images[0]]
|
||||
val_images = [images[1]]
|
||||
num_groups = len(group_keys)
|
||||
if num_groups == 1:
|
||||
train_groups = group_keys
|
||||
val_groups = group_keys
|
||||
elif num_groups == 2:
|
||||
train_groups = [group_keys[0]]
|
||||
val_groups = [group_keys[1]]
|
||||
else:
|
||||
split_idx = max(1, int(num_images * split_ratio))
|
||||
split_idx = min(split_idx, num_images - 1)
|
||||
train_images = images[:split_idx]
|
||||
val_images = images[split_idx:]
|
||||
split_idx = max(1, int(num_groups * split_ratio))
|
||||
split_idx = min(split_idx, num_groups - 1)
|
||||
train_groups = group_keys[:split_idx]
|
||||
val_groups = group_keys[split_idx:]
|
||||
|
||||
train_images = [img for key in train_groups for img in groups[key]]
|
||||
val_images = [img for key in val_groups for img in groups[key]]
|
||||
|
||||
for img in train_images:
|
||||
shutil.copy(img, class_train_dir / img.name)
|
||||
@@ -145,7 +169,7 @@ def split_dataset(src_dir: Path, dest_dir: Path, split_ratio: float = 0.8, seed:
|
||||
|
||||
print(
|
||||
f" Class '{class_name}': {len(train_images)} train, "
|
||||
f"{len(val_images)} val (total: {num_images})"
|
||||
f"{len(val_images)} val (from {num_groups} source photos, {num_images} files total)"
|
||||
)
|
||||
|
||||
print(f"Dataset split completed: {total_train} train images, {total_val} validation images.")
|
||||
|
||||
Reference in new issue
Block a user