From f052c57aa880efe8096bfa8da66f8e8fec7b2069 Mon Sep 17 00:00:00 2001 From: James <1561862923@qq.com> Date: Sat, 15 Aug 2026 18:23:40 +0800 Subject: [PATCH] fix: 1 --- testing/test_trigger_selective_training.py | 42 ++++++++++++++++++++++ toolkit/data_loader.py | 22 +++++++++++- 2 files changed, 63 insertions(+), 1 deletion(-) diff --git a/testing/test_trigger_selective_training.py b/testing/test_trigger_selective_training.py index e13db447..0faccdad 100644 --- a/testing/test_trigger_selective_training.py +++ b/testing/test_trigger_selective_training.py @@ -257,6 +257,48 @@ class TriggerSelectiveTrainingTest(unittest.TestCase): self.assertEqual(item['item_id'], os.path.join('nested', 'item.png')) self.assertEqual(item['sources']['natural']['caption'], 'natural [trigger] caption') + def test_caption_source_orphan_check_uses_complete_pre_split_image_set(self): + with tempfile.TemporaryDirectory() as temp_dir: + main_root = os.path.join(temp_dir, 'main') + mirror_root = os.path.join(temp_dir, 'mirror') + os.makedirs(main_root) + os.makedirs(mirror_root) + train_image = os.path.join(main_root, 'train.png') + heldout_image = os.path.join(main_root, 'heldout.png') + for image_path in (train_image, heldout_image): + open(image_path, 'wb').close() + with open(os.path.splitext(image_path)[0] + '.json', 'w', encoding='utf-8') as handle: + json.dump({'caption': 'structured [trigger]'}, handle) + mirror_image = os.path.join(mirror_root, os.path.basename(image_path)) + open(mirror_image, 'wb').close() + with open(os.path.splitext(mirror_image)[0] + '.txt', 'w', encoding='utf-8') as handle: + handle.write('natural [trigger]') + config = TriggerSelectiveTrainingConfig( + caption_sources={ + 'enabled': True, + 'sources': [ + {'name': 'structured', 'use_main_dataset': True, 'caption_ext': '.json', 'format': 'json'}, + {'name': 'natural', 'path': mirror_root, 'caption_ext': '.txt', 'format': 'text'}, + ], + }, + ) + result = discover_tst_caption_sources( + main_root, + [train_image], + config.caption_sources, + complete_file_list_for_orphan_check=[train_image, heldout_image], + ) + self.assertEqual(set(result), {os.path.abspath(train_image)}) + extra_image = os.path.join(mirror_root, 'real_orphan.png') + open(extra_image, 'wb').close() + with self.assertRaisesRegex(ValueError, 'orphan mirror image'): + discover_tst_caption_sources( + main_root, + [train_image], + config.caption_sources, + complete_file_list_for_orphan_check=[train_image, heldout_image], + ) + def test_json_only_main_caption_source_is_supported(self): config = TriggerSelectiveTrainingConfig( enabled=True, diff --git a/toolkit/data_loader.py b/toolkit/data_loader.py index c450a27f..b1f909df 100644 --- a/toolkit/data_loader.py +++ b/toolkit/data_loader.py @@ -41,11 +41,20 @@ video_extensions = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.f audio_extensions = ['.mp3', '.wav', '.flac', '.aac', '.ogg', '.m4a'] -def discover_tst_caption_sources(dataset_root, file_list, caption_sources): +def discover_tst_caption_sources( + dataset_root, + file_list, + caption_sources, + complete_file_list_for_orphan_check=None, +): if not caption_sources or not caption_sources.enabled: return {} dataset_root = os.path.abspath(dataset_root) unique_files = list(dict.fromkeys(os.path.abspath(path) for path in file_list)) + complete_files = list(dict.fromkeys( + os.path.abspath(path) + for path in (complete_file_list_for_orphan_check or unique_files) + )) source_map = {} trigger_counts = {source.name: {} for source in caption_sources.sources} expected_mirror_images = set() @@ -86,6 +95,15 @@ def discover_tst_caption_sources(dataset_root, file_list, caption_sources): if source.use_main_dataset: continue source_root = os.path.abspath(source.path) + for image_path in complete_files: + relative_id = os.path.normpath(os.path.relpath(image_path, dataset_root)) + relative_stem = os.path.splitext(relative_id)[0] + expected_mirror_images.add(os.path.normcase(os.path.abspath( + os.path.join(source_root, relative_id) + ))) + expected_mirror_captions.add(os.path.normcase(os.path.abspath( + os.path.join(source_root, relative_stem + source.caption_ext) + ))) actual_images = set() actual_captions = set() for root, dirs, files in os.walk(source_root): @@ -521,6 +539,7 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin file_list = [x for x in file_list if not os.path.basename(os.path.dirname(x)) == "_controls"] tst_dataset_root = self.dataset_path if os.path.isdir(self.dataset_path) else os.path.dirname(self.dataset_path) + complete_file_list_for_orphan_check = list(file_list) split_manifest_path = getattr(self.dataset_config, 'trigger_data_split_manifest', None) split_name = getattr(self.dataset_config, 'trigger_data_split_name', None) if split_manifest_path is not None: @@ -545,6 +564,7 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin tst_dataset_root, file_list, self.dataset_config.trigger_selective_caption_sources, + complete_file_list_for_orphan_check=complete_file_list_for_orphan_check, ) if self.dataset_config.num_repeats > 1: