From 41157b460c671745e60fb4ea6359080aab4f1216 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sat, 6 Jun 2026 08:32:24 -0600 Subject: [PATCH] Added ability to set the caption extention in dataset viewer, captioner, and trainer so one dataset can have multiple caption styles in different files with different extensions. Added dataset caption template for a blank ideogram 4 formatted template. --- toolkit/config_modules.py | 5 +- toolkit/data_loader.py | 5 - toolkit/data_transfer_object/data_loader.py | 2 - toolkit/dataloader_mixins.py | 195 +------------------- ui/src/app/api/caption/get/route.ts | 7 +- ui/src/app/api/caption/getBatch/route.ts | 5 +- ui/src/app/api/img/caption/route.ts | 7 +- ui/src/app/datasets/[datasetName]/page.tsx | 18 ++ ui/src/app/jobs/new/SimpleJob.tsx | 13 ++ ui/src/components/AutoCaptionButton.tsx | 13 +- ui/src/components/CaptionDatasetModal.tsx | 9 +- ui/src/components/CaptionSimpleJob.tsx | 14 ++ ui/src/components/DatasetImageCard.tsx | 17 +- ui/src/components/DatasetImageViewer.tsx | 40 +++- ui/src/components/formInputs.tsx | 6 +- ui/src/helpers/captionJobConfig.ts | 1 + ui/src/helpers/datasetTemplates.ts | 21 +++ ui/src/hooks/useCaptionBatch.tsx | 110 ++++++----- ui/src/types.ts | 1 + 19 files changed, 219 insertions(+), 270 deletions(-) create mode 100644 ui/src/helpers/datasetTemplates.ts diff --git a/toolkit/config_modules.py b/toolkit/config_modules.py index 3155e7b..a42c4a5 100644 --- a/toolkit/config_modules.py +++ b/toolkit/config_modules.py @@ -944,8 +944,9 @@ class DatasetConfig: None) # path where matching unconditional images are located self.invert_mask: bool = kwargs.get('invert_mask', False) # invert mask self.mask_min_value: float = kwargs.get('mask_min_value', 0.0) # min value for . 0 - 1 - self.poi: Union[str, None] = kwargs.get('poi', - None) # if one is set and in json data, will be used as auto crop scale point of interes + self.poi: Union[str, None] = kwargs.get('poi', None) + if self.poi is not None: + raise ValueError("poi is deprecated and is no longer supported") self.use_short_captions: bool = kwargs.get('use_short_captions', False) # if true, will use 'caption_short' from json self.num_repeats: int = kwargs.get('num_repeats', 1) # number of times to repeat dataset # cache latents will store them in memory diff --git a/toolkit/data_loader.py b/toolkit/data_loader.py index 9fd5f86..e412e51 100644 --- a/toolkit/data_loader.py +++ b/toolkit/data_loader.py @@ -604,11 +604,6 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin if self.is_generating_controls: # always do this last self.setup_controls() - else: - if self.dataset_config.poi is not None: - # handle cropping to a specific point of interest - # setup buckets every epoch - self.setup_buckets(quiet=True) self.epoch_num += 1 def __len__(self): diff --git a/toolkit/data_transfer_object/data_loader.py b/toolkit/data_transfer_object/data_loader.py index 36ed176..681be45 100644 --- a/toolkit/data_transfer_object/data_loader.py +++ b/toolkit/data_transfer_object/data_loader.py @@ -15,7 +15,6 @@ from toolkit.dataloader_mixins import ( LatentCachingFileItemDTOMixin, ControlFileItemDTOMixin, ArgBreakMixin, - PoiFileItemDTOMixin, MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, UnconditionalFileItemDTOMixin, @@ -51,7 +50,6 @@ class FileItemDTO( MaskFileItemDTOMixin, AugmentationFileItemDTOMixin, UnconditionalFileItemDTOMixin, - PoiFileItemDTOMixin, ArgBreakMixin, ): def __init__(self, *args, **kwargs): diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index 02d8eba..d9ce629 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -150,15 +150,9 @@ class CaptionMixin: if os.path.exists(prompt_path): with open(prompt_path, 'r', encoding='utf-8') as f: prompt = f.read() - # check if is json - if prompt_path.endswith('.json'): - prompt = json.loads(prompt) - if 'caption' in prompt: - prompt = prompt['caption'] - prompt = clean_caption(prompt) elif os.path.exists(default_prompt_path_with_ext): - with open(default_prompt_path, 'r', encoding='utf-8') as f: + with open(default_prompt_path_with_ext, 'r', encoding='utf-8') as f: prompt = f.read() prompt = clean_caption(prompt) elif os.path.exists(default_prompt_path): @@ -217,7 +211,7 @@ class BucketsMixin: if not hasattr(self, 'dataset_config'): raise Exception(f'dataset_config not found on class instance {self.__class__.__name__}') - if self.epoch_num > 0 and self.dataset_config.poi is None: + if self.epoch_num > 0: # no need to rebuild buckets for now # todo handle random cropping for buckets return @@ -240,10 +234,6 @@ class BucketsMixin: width = int(file_item.width * file_item.dataset_config.scale) height = int(file_item.height * file_item.dataset_config.scale) - did_process_poi = False - if file_item.has_point_of_interest: - # Attempt to process the poi if we can. It wont process if the image is smaller than the resolution - did_process_poi = file_item.setup_poi_bucket() if self.dataset_config.square_crop: # we scale first so smallest size matches resolution scale_factor_x = resolution / width @@ -259,7 +249,7 @@ class BucketsMixin: else: file_item.crop_x = 0 file_item.crop_y = int(file_item.scale_to_height / 2 - resolution / 2) - elif not did_process_poi: + else: bucket_resolution = get_bucket_for_image_size( width, height, resolution=resolution, @@ -348,22 +338,6 @@ class CaptionProcessingDTOMixin: with open(prompt_path, 'r', encoding='utf-8') as f: prompt = f.read() short_caption = None - if prompt_path.endswith('.json'): - # replace any line endings with commas for \n \r \r\n - prompt = prompt.replace('\r\n', ' ') - prompt = prompt.replace('\n', ' ') - prompt = prompt.replace('\r', ' ') - - prompt_json = json.loads(prompt) - if 'caption' in prompt_json: - prompt = prompt_json['caption'] - if 'caption_short' in prompt_json: - short_caption = prompt_json['caption_short'] - if self.dataset_config.use_short_captions: - prompt = short_caption - if 'extra_values' in prompt_json: - self.extra_values = prompt_json['extra_values'] - prompt = clean_caption(prompt) if short_caption is not None: short_caption = clean_caption(short_caption) @@ -414,10 +388,6 @@ class CaptionProcessingDTOMixin: # get tokens token_list = raw_caption.split(',') - # trim whitespace - token_list = [x.strip() for x in token_list] - # remove empty strings - token_list = [x for x in token_list if x] # handle token dropout if self.dataset_config.token_dropout_rate > 0 and not short_caption and not self.dataset_config.cache_text_embeddings: @@ -461,10 +431,6 @@ class CaptionProcessingDTOMixin: if self.dataset_config.shuffle_tokens: # shuffle again token_list = caption.split(',') - # trim whitespace - token_list = [x.strip() for x in token_list] - # remove empty strings - token_list = [x for x in token_list if x] random.shuffle(token_list) caption = ', '.join(token_list) if caption == '': @@ -1607,161 +1573,6 @@ class UnconditionalFileItemDTOMixin: self.unconditional_tensor = None self.unconditional_latent = None - -class PoiFileItemDTOMixin: - # Point of interest bounding box. Allows for dynamic cropping without cropping out the main subject - # items in the poi will always be inside the image when random cropping - def __init__(self: 'FileItemDTO', *args, **kwargs): - if hasattr(super(), '__init__'): - super().__init__(*args, **kwargs) - # poi is a name of the box point of interest in the caption json file - dataset_config = kwargs.get('dataset_config', None) - path = kwargs.get('path', None) - self.poi: Union[str, None] = dataset_config.poi - self.has_point_of_interest = self.poi is not None - self.poi_x: Union[int, None] = None - self.poi_y: Union[int, None] = None - self.poi_width: Union[int, None] = None - self.poi_height: Union[int, None] = None - - if self.poi is not None: - # make sure latent caching is off - if dataset_config.cache_latents or dataset_config.cache_latents_to_disk: - raise Exception( - f"Error: poi is not supported when caching latents. Please set cache_latents and cache_latents_to_disk to False in the dataset config" - ) - # make sure we are loading through json - if dataset_config.caption_ext != 'json': - raise Exception( - f"Error: poi is only supported when using json captions. Please set caption_ext to json in the dataset config" - ) - self.poi = self.poi.strip() - # get the caption path - file_path_no_ext = os.path.splitext(path)[0] - caption_path = file_path_no_ext + '.json' - if not os.path.exists(caption_path): - raise Exception(f"Error: caption file not found for poi: {caption_path}") - with open(caption_path, 'r', encoding='utf-8') as f: - json_data = json.load(f) - if 'poi' not in json_data: - print_acc(f"Warning: poi not found in caption file: {caption_path}") - if self.poi not in json_data['poi']: - print_acc(f"Warning: poi not found in caption file: {caption_path}") - # poi has, x, y, width, height - # do full image if no poi - self.poi_x = 0 - self.poi_y = 0 - self.poi_width = self.width - self.poi_height = self.height - try: - if self.poi in json_data['poi']: - poi = json_data['poi'][self.poi] - self.poi_x = int(poi['x']) - self.poi_y = int(poi['y']) - self.poi_width = int(poi['width']) - self.poi_height = int(poi['height']) - except Exception as e: - pass - - # handle flipping - if kwargs.get('flip_x', False): - # flip the poi - self.poi_x = self.width - self.poi_x - self.poi_width - if kwargs.get('flip_y', False): - # flip the poi - self.poi_y = self.height - self.poi_y - self.poi_height - - def setup_poi_bucket(self: 'FileItemDTO'): - initial_width = int(self.width * self.dataset_config.scale) - initial_height = int(self.height * self.dataset_config.scale) - # we are using poi, so we need to calculate the bucket based on the poi - - # if img resolution is less than dataset resolution, just return and let the normal bucketing happen - img_resolution = get_resolution(initial_width, initial_height) - if img_resolution <= self.dataset_config.resolution: - return False # will trigger normal bucketing - - bucket_tolerance = self.dataset_config.bucket_tolerance - poi_x = int(self.poi_x * self.dataset_config.scale) - poi_y = int(self.poi_y * self.dataset_config.scale) - poi_width = int(self.poi_width * self.dataset_config.scale) - poi_height = int(self.poi_height * self.dataset_config.scale) - - # loop to keep expanding until we are at the proper resolution. This is not ideal, we can probably handle it better - num_loops = 0 - while True: - # crop left - if poi_x > 0: - poi_x = random.randint(0, poi_x) - else: - poi_x = 0 - - # crop right - cr_min = poi_x + poi_width - if cr_min < initial_width: - crop_right = random.randint(poi_x + poi_width, initial_width) - else: - crop_right = initial_width - - poi_width = crop_right - poi_x - - if poi_y > 0: - poi_y = random.randint(0, poi_y) - else: - poi_y = 0 - - if poi_y + poi_height < initial_height: - crop_bottom = random.randint(poi_y + poi_height, initial_height) - else: - crop_bottom = initial_height - - poi_height = crop_bottom - poi_y - try: - # now we have our random crop, but it may be smaller than resolution. Check and expand if needed - current_resolution = get_resolution(poi_width, poi_height) - except Exception as e: - print_acc(f"Error: {e}") - print_acc(f"Error getting resolution: {self.path}") - raise e - return False - if current_resolution >= self.dataset_config.resolution: - # We can break now - break - else: - num_loops += 1 - if num_loops > 100: - print_acc( - f"Warning: poi bucketing looped too many times. This should not happen. Please report this issue.") - return False - - new_width = poi_width - new_height = poi_height - - bucket_resolution = get_bucket_for_image_size( - new_width, new_height, - resolution=self.dataset_config.resolution, - divisibility=bucket_tolerance - ) - - width_scale_factor = bucket_resolution["width"] / new_width - height_scale_factor = bucket_resolution["height"] / new_height - # Use the maximum of the scale factors to ensure both dimensions are scaled above the bucket resolution - max_scale_factor = max(width_scale_factor, height_scale_factor) - - self.scale_to_width = math.ceil(initial_width * max_scale_factor) - self.scale_to_height = math.ceil(initial_height * max_scale_factor) - self.crop_width = bucket_resolution['width'] - self.crop_height = bucket_resolution['height'] - self.crop_x = int(poi_x * max_scale_factor) - self.crop_y = int(poi_y * max_scale_factor) - - if self.scale_to_width < self.crop_x + self.crop_width or self.scale_to_height < self.crop_y + self.crop_height: - # todo look into this. This still happens sometimes - print_acc('size mismatch') - - return True - - class ArgBreakMixin: # just stops super calls form hitting object def __init__(self, *args, **kwargs): diff --git a/ui/src/app/api/caption/get/route.ts b/ui/src/app/api/caption/get/route.ts index e1b202b..887e70f 100644 --- a/ui/src/app/api/caption/get/route.ts +++ b/ui/src/app/api/caption/get/route.ts @@ -22,15 +22,16 @@ export async function POST(request: NextRequest) { return new NextResponse(null, { status: 499 }); } - const { imgPath } = body; + const { imgPath, ext } = body; console.log('Received POST request for caption:', imgPath); try { // Decode the path const filepath = imgPath; console.log('Decoded image path:', filepath); - // caption name is the filepath without extension but with .txt - const captionPath = filepath.replace(/\.[^/.]+$/, '') + '.txt'; + // caption name is the filepath without extension but with the caption extension (default txt) + const captionExt = ((ext || 'txt') as string).replace(/^\.+/, '').trim() || 'txt'; + const captionPath = filepath.replace(/\.[^/.]+$/, '') + '.' + captionExt; // Get allowed directories const allowedDir = await getDatasetsRoot(); diff --git a/ui/src/app/api/caption/getBatch/route.ts b/ui/src/app/api/caption/getBatch/route.ts index 281c16a..5968fbe 100644 --- a/ui/src/app/api/caption/getBatch/route.ts +++ b/ui/src/app/api/caption/getBatch/route.ts @@ -21,11 +21,12 @@ export async function POST(request: NextRequest) { return new NextResponse(null, { status: 499 }); } - const { imgPaths } = body as { imgPaths?: string[] }; + const { imgPaths, ext } = body as { imgPaths?: string[]; ext?: string }; if (!Array.isArray(imgPaths)) { return NextResponse.json({ error: 'imgPaths must be an array' }, { status: 400 }); } + const captionExt = ((ext || 'txt') as string).replace(/^\.+/, '').trim() || 'txt'; const allowedDir = await getDatasetsRoot(); const captions: Record = {}; @@ -33,7 +34,7 @@ export async function POST(request: NextRequest) { if (typeof imgPath !== 'string') continue; if (!isUnderRoot(imgPath, allowedDir)) continue; - const captionPath = imgPath.replace(/\.[^/.]+$/, '') + '.txt'; + const captionPath = imgPath.replace(/\.[^/.]+$/, '') + '.' + captionExt; try { captions[imgPath] = fs.existsSync(captionPath) ? fs.readFileSync(captionPath, 'utf-8') : ''; } catch { diff --git a/ui/src/app/api/img/caption/route.ts b/ui/src/app/api/img/caption/route.ts index df4235f..7aa82e9 100644 --- a/ui/src/app/api/img/caption/route.ts +++ b/ui/src/app/api/img/caption/route.ts @@ -5,7 +5,7 @@ import { getDatasetsRoot } from '@/server/settings'; export async function POST(request: Request) { try { const body = await request.json(); - const { imgPath, caption } = body; + const { imgPath, caption, ext } = body; let datasetsPath = await getDatasetsRoot(); // make sure the dataset path is in the image path if (!imgPath.startsWith(datasetsPath)) { @@ -17,8 +17,9 @@ export async function POST(request: Request) { return NextResponse.json({ error: 'Image does not exist' }, { status: 404 }); } - // check for caption - const captionPath = imgPath.replace(/\.[^/.]+$/, '') + '.txt'; + // check for caption (default extension txt) + const captionExt = ((ext || 'txt') as string).replace(/^\.+/, '').trim() || 'txt'; + const captionPath = imgPath.replace(/\.[^/.]+$/, '') + '.' + captionExt; // save caption to file fs.writeFileSync(captionPath, caption); diff --git a/ui/src/app/datasets/[datasetName]/page.tsx b/ui/src/app/datasets/[datasetName]/page.tsx index 9dc7575..29dbaa6 100644 --- a/ui/src/app/datasets/[datasetName]/page.tsx +++ b/ui/src/app/datasets/[datasetName]/page.tsx @@ -13,6 +13,7 @@ import { apiClient } from '@/utils/api'; import useSettings from '@/hooks/useSettings'; import { pathJoin } from '@/utils/basic'; import AutoCaptionButton from '@/components/AutoCaptionButton'; +import { CreatableSelectInput } from '@/components/formInputs'; export default function DatasetPage({ params }: { params: { datasetName: string } }) { const [imgList, setImgList] = useState<{ img_path: string }[]>([]); @@ -22,6 +23,7 @@ export default function DatasetPage({ params }: { params: { datasetName: string const [status, setStatus] = useState<'idle' | 'loading' | 'success' | 'error'>('idle'); const { settings, isSettingsLoaded } = useSettings(); const [selectedImgPath, setSelectedImgPath] = useState(null); + const [captionExt, setCaptionExt] = useState('txt'); const [captionRefreshKeys, setCaptionRefreshKeys] = useState>({}); const [scrollParent, setScrollParent] = useState(null); const scrollParentCallback = useCallback((el: HTMLDivElement | null) => setScrollParent(el), []); @@ -118,9 +120,23 @@ export default function DatasetPage({ params }: { params: { datasetName: string
+
+ + setCaptionExt(value)} + options={[ + { value: 'txt', label: 'txt' }, + { value: 'json', label: 'json' }, + { value: 'caption', label: 'caption' }, + ]} + /> +
+ {isCaptionLoaded && caption.trim() === '' && ( + + )} {isIdeogram ? ( { )}
-
+