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.

This commit is contained in:
Jaret Burkett
2026-06-06 08:32:24 -06:00
parent 10cdeb394e
commit 41157b460c
19 changed files with 219 additions and 270 deletions

View File

@@ -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

View File

@@ -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):

View File

@@ -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):

View File

@@ -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):

View File

@@ -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();

View File

@@ -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<string, string> = {};
@@ -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 {

View File

@@ -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);

View File

@@ -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<string | null>(null);
const [captionExt, setCaptionExt] = useState<string>('txt');
const [captionRefreshKeys, setCaptionRefreshKeys] = useState<Record<string, number>>({});
const [scrollParent, setScrollParent] = useState<HTMLDivElement | null>(null);
const scrollParentCallback = useCallback((el: HTMLDivElement | null) => setScrollParent(el), []);
@@ -118,9 +120,23 @@ export default function DatasetPage({ params }: { params: { datasetName: string
</div>
<div className="flex-1"></div>
<div className="flex-shrink-0 flex items-center gap-1 sm:gap-2">
<div className="flex items-center gap-1">
<label className="text-xs text-gray-400 hidden sm:inline whitespace-nowrap">Caption ext</label>
<CreatableSelectInput
className="w-44"
value={captionExt}
onChange={value => setCaptionExt(value)}
options={[
{ value: 'txt', label: 'txt' },
{ value: 'json', label: 'json' },
{ value: 'caption', label: 'caption' },
]}
/>
</div>
<AutoCaptionButton
datasetPath={`${pathJoin(settings.DATASETS_FOLDER, datasetName)}`}
setIsAutoCaptioning={setIsAutoCaptioning}
captionExt={captionExt}
/>
<Button
className="text-white bg-slate-600 px-2 sm:px-3 py-1 rounded-md text-sm sm:text-base whitespace-nowrap"
@@ -151,6 +167,7 @@ export default function DatasetPage({ params }: { params: { datasetName: string
onImageClick={() => setSelectedImgPath(img.img_path)}
captionRefreshKey={captionRefreshKeys[img.img_path] || 0}
observerRoot={scrollParent}
captionExt={captionExt}
/>
);
}}
@@ -165,6 +182,7 @@ export default function DatasetPage({ params }: { params: { datasetName: string
onChange={setSelectedImgPath}
refreshImages={() => refreshImageList(datasetName)}
onCaptionSaved={path => setCaptionRefreshKeys(prev => ({ ...prev, [path]: (prev[path] || 0) + 1 }))}
captionExt={captionExt}
/>
</>
);

View File

@@ -20,6 +20,7 @@ import {
FormGroup,
NumberInput,
SliderInput,
CreatableSelectInput,
} from '@/components/formInputs';
import Card from '@/components/Card';
import { X, Copy, Wand2, SquareDashed } from 'lucide-react';
@@ -959,6 +960,18 @@ export default function SimpleJob({
min={0}
required
/>
<CreatableSelectInput
label="Caption Extension"
className="pt-2"
value={dataset.caption_ext || 'txt'}
onChange={value => setJobConfig(value, `config.process[0].datasets[${i}].caption_ext`)}
options={[
{ value: 'txt', label: 'txt' },
{ value: 'json', label: 'json' },
{ value: 'caption', label: 'caption' },
]}
/>
{modelArch?.additionalSections?.includes('datasets.num_frames') && !dataset.auto_frame_count && (
<NumberInput
label="Num Frames"

View File

@@ -8,9 +8,10 @@ import { Loader2 } from 'lucide-react';
type AutoCaptionButtonProps = {
datasetPath: string;
setIsAutoCaptioning?: (isAutoCaptioning: boolean) => void;
captionExt?: string;
};
export default function AutoCaptionButton({ datasetPath, setIsAutoCaptioning }: AutoCaptionButtonProps) {
export default function AutoCaptionButton({ datasetPath, setIsAutoCaptioning, captionExt }: AutoCaptionButtonProps) {
const { job, status, refreshJob } = useJobByRef(datasetPath, 5000);
useEffect(() => {
if (setIsAutoCaptioning) {
@@ -34,9 +35,13 @@ export default function AutoCaptionButton({ datasetPath, setIsAutoCaptioning }:
<Button
className="text-white bg-blue-600 px-2 sm:px-3 py-1 rounded-md mr-1 sm:mr-2 text-sm sm:text-base whitespace-nowrap"
onClick={() =>
openCaptionDatasetModal(datasetPath, () => {
refreshJob();
})
openCaptionDatasetModal(
datasetPath,
() => {
refreshJob();
},
{ defaultCaptionExt: captionExt },
)
}
>
<span className="hidden sm:inline">Auto Caption</span>

View File

@@ -22,6 +22,7 @@ export interface CaptionDatasetModalState {
datasetPath: string;
jobId?: string | null;
cloneId?: string | null;
defaultCaptionExt?: string;
onClose?: () => void;
}
@@ -30,13 +31,14 @@ export const captionDatasetModalState = createGlobalState<CaptionDatasetModalSta
export const openCaptionDatasetModal = (
datasetPath: string,
onClose?: () => void,
options?: { jobId?: string | null; cloneId?: string | null },
options?: { jobId?: string | null; cloneId?: string | null; defaultCaptionExt?: string },
) => {
captionDatasetModalState.set({
datasetPath,
onClose,
jobId: options?.jobId ?? null,
cloneId: options?.cloneId ?? null,
defaultCaptionExt: options?.defaultCaptionExt,
});
};
@@ -64,6 +66,11 @@ export const CaptionDatasetModal: React.FC = () => {
if (modalInfo?.datasetPath) {
setJobConfig(modalInfo.datasetPath, 'config.process[0].caption.path_to_caption');
}
// default the caption extension to the current header selection. Editing it
// here only changes this job, not the header (separate state).
if (modalInfo?.defaultCaptionExt) {
setJobConfig(modalInfo.defaultCaptionExt, 'config.process[0].caption.caption_extension');
}
}, [modalInfo]);
// clone existing caption job

View File

@@ -121,6 +121,20 @@ const CaptionSimpleJob: React.FC<Props> = ({ jobConfig, setJobConfig, gpuIDs, se
}}
options={quantizationOptions}
/>
<div className="mt-4">
<CreatableSelectInput
label="Caption Extension"
value={jobConfig.config.process[0].caption.caption_extension || 'txt'}
onChange={value => {
setJobConfig(value, 'config.process[0].caption.caption_extension');
}}
options={[
{ value: 'txt', label: 'txt' },
{ value: 'json', label: 'json' },
{ value: 'caption', label: 'caption' },
]}
/>
</div>
{additionalSections.includes('caption.max_res') && (
<div className="mt-4">
<SelectInput

View File

@@ -18,6 +18,7 @@ interface DatasetImageCardProps {
captionRefreshKey?: number;
observerRoot?: Element | null;
rootMargin?: string;
captionExt?: string;
}
const DatasetImageCard: React.FC<DatasetImageCardProps> = ({
@@ -31,6 +32,7 @@ const DatasetImageCard: React.FC<DatasetImageCardProps> = ({
captionRefreshKey = 0,
observerRoot = null,
rootMargin = '200px 0px',
captionExt = 'txt',
}) => {
const [loaded, setLoaded] = useState<boolean>(false);
const [showAudioPlayer, setShowAudioPlayer] = useState(true);
@@ -110,6 +112,7 @@ const DatasetImageCard: React.FC<DatasetImageCardProps> = ({
const { caption: fetchedCaption, isLoaded: isCaptionLoaded } = useCaptionBatch(
isVisible ? imageUrl : null,
combinedRefreshKey,
captionExt,
);
const [caption, setCaption] = useState<string>('');
@@ -138,10 +141,10 @@ const DatasetImageCard: React.FC<DatasetImageCardProps> = ({
return;
}
apiClient
.post('/api/img/caption', { imgPath: imageUrl, caption: trimmedCaption })
.post('/api/img/caption', { imgPath: imageUrl, caption: trimmedCaption, ext: captionExt })
.then(() => {
setSavedCaption(trimmedCaption);
setCachedCaption(imageUrl, trimmedCaption);
setCachedCaption(imageUrl, trimmedCaption, captionExt);
dirtyRef.current = false;
})
.catch(error => {
@@ -150,19 +153,19 @@ const DatasetImageCard: React.FC<DatasetImageCardProps> = ({
};
// Save any pending edit if the card unmounts (e.g. scrolled out of the virtualized window).
const latestRef = useRef({ caption, savedCaption, imageUrl });
const latestRef = useRef({ caption, savedCaption, imageUrl, captionExt });
useEffect(() => {
latestRef.current = { caption, savedCaption, imageUrl };
latestRef.current = { caption, savedCaption, imageUrl, captionExt };
});
useEffect(() => {
return () => {
if (!dirtyRef.current) return;
const { caption: c, savedCaption: s, imageUrl: url } = latestRef.current;
const { caption: c, savedCaption: s, imageUrl: url, captionExt: ext } = latestRef.current;
const trimmed = c.trim();
if (trimmed === s) return;
apiClient
.post('/api/img/caption', { imgPath: url, caption: trimmed })
.then(() => setCachedCaption(url, trimmed))
.post('/api/img/caption', { imgPath: url, caption: trimmed, ext })
.then(() => setCachedCaption(url, trimmed, ext))
.catch(err => console.error('Error saving caption on unmount:', err));
};
}, []);

View File

@@ -12,6 +12,7 @@ import AudioPlayer from './AudioPlayer';
import { TransformWrapper, TransformComponent } from 'react-zoom-pan-pinch';
import { BoundingBoxEditor, parseBoundingBoxes, extractBoxes } from './BoundingBoxOverlay';
import IdeogramCaptionSidebar, { isIdeogramCaption } from './IdeogramCaptionSidebar';
import datasetTemplates from '@/helpers/datasetTemplates';
function safeParse(text: string): any {
try {
@@ -27,9 +28,17 @@ interface Props {
onChange: (nextPath: string | null) => void; // parent setter
refreshImages?: () => void;
onCaptionSaved?: (imgPath: string, caption: string) => void;
captionExt?: string;
}
export default function DatasetImageViewer({ imgPath, imageList, onChange, refreshImages, onCaptionSaved }: Props) {
export default function DatasetImageViewer({
imgPath,
imageList,
onChange,
refreshImages,
onCaptionSaved,
captionExt = 'txt',
}: Props) {
const [mounted, setMounted] = useState(false);
const [isOpen, setIsOpen] = useState(Boolean(imgPath));
const [caption, setCaption] = useState<string>('');
@@ -98,7 +107,7 @@ export default function DatasetImageViewer({ imgPath, imageList, onChange, refre
const trimmed = value.trim();
if (trimmed === prevSaved.trim()) return;
apiClient
.post('/api/img/caption', { imgPath: path, caption: trimmed })
.post('/api/img/caption', { imgPath: path, caption: trimmed, ext: captionExt })
.then(() => {
if (currentImgPathRef.current === path) {
setSavedCaption(trimmed);
@@ -109,7 +118,7 @@ export default function DatasetImageViewer({ imgPath, imageList, onChange, refre
console.error('Error saving caption:', error);
});
},
[onCaptionSaved],
[onCaptionSaved, captionExt],
);
// Stable handle to the latest saveCaptionForPath so the fetch effect doesn't
@@ -151,7 +160,11 @@ export default function DatasetImageViewer({ imgPath, imageList, onChange, refre
// transformResponse identity: keep the caption as a raw string. Axios's
// default parses any JSON-looking body into an object (our bbox captions
// are JSON), which would render as "[object Object]".
.post('/api/caption/get', { imgPath }, { signal: controller.signal, transformResponse: [d => d] })
.post(
'/api/caption/get',
{ imgPath, ext: captionExt },
{ signal: controller.signal, transformResponse: [d => d] },
)
.then(res => res.data)
.then(data => {
if (controller.signal.aborted) return;
@@ -169,7 +182,7 @@ export default function DatasetImageViewer({ imgPath, imageList, onChange, refre
return () => {
controller.abort();
};
}, [imgPath]);
}, [imgPath, captionExt]);
// Save any pending caption when the viewer fully unmounts
useEffect(() => {
@@ -509,6 +522,23 @@ export default function DatasetImageViewer({ imgPath, imageList, onChange, refre
{currentIndex >= 0 ? `${currentIndex + 1} / ${imageList.length}` : ''}
</div>
</div>
{isCaptionLoaded && caption.trim() === '' && (
<select
className="w-full bg-gray-900 border border-gray-700 text-gray-100 text-sm rounded p-2 outline-none focus:ring-0 focus:outline-none"
value=""
onChange={e => {
const template = datasetTemplates[e.target.value];
if (template) setCaption(template.trim());
}}
>
<option value="">Templates...</option>
{Object.keys(datasetTemplates).map(key => (
<option key={key} value={key}>
{key}
</option>
))}
</select>
)}
{isIdeogram ? (
<IdeogramCaptionSidebar
caption={caption}

View File

@@ -363,13 +363,15 @@ export const CreatableSelectInput = (props: CreatableSelectInputProps) => {
</label>
)}
<div className="flex gap-2">
<div className={isCustom ? 'w-1/3' : 'w-full'}>
<div className={isCustom ? 'w-20 shrink-0' : 'w-full'}>
<Select
value={selectedOption}
options={selectOptions}
isDisabled={props.disabled}
className="aitk-react-select-container"
classNamePrefix="aitk-react-select"
menuPosition="fixed"
menuPlacement="auto"
formatOptionLabel={(option: unknown) => {
const opt = option as SelectOption;
return opt.value === CUSTOM_SELECT_VALUE ? (
@@ -397,7 +399,7 @@ export const CreatableSelectInput = (props: CreatableSelectInputProps) => {
type="text"
value={value}
onChange={e => onChange(e.target.value)}
className={`${inputClasses} w-2/3`}
className={`${inputClasses} flex-1 min-w-0`}
placeholder={props.placeholder ?? 'Enter custom value'}
disabled={props.disabled}
autoFocus

View File

@@ -21,6 +21,7 @@ export const defaultCaptionJobConfig: CaptionJobConfig = {
extensions: ['mp3', 'wav', 'flac', 'ogg'],
path_to_caption: '',
recaption: false,
caption_extension: 'txt',
},
},
],

View File

@@ -0,0 +1,21 @@
const datasetTemplates: { [key: string]: string } = {
ideogram4: `
{
"high_level_description": "",
"style_description": {
"aesthetics": "",
"lighting": "",
"photo": "",
"medium": "",
"color_palette": []
},
"compositional_deconstruction": {
"background": "",
"elements": [
]
}
}
`
};
export default datasetTemplates;

View File

@@ -4,13 +4,24 @@ import { apiClient } from '@/utils/api';
// Module-level batcher: many cards mount at once when the virtualized grid scrolls;
// instead of N HTTP requests, queue paths and flush them as a single batch.
// Entries are keyed by extension + path so different caption extensions don't
// collide in the cache or the pending batch.
type Resolver = { resolve: (caption: string) => void; reject: (err: unknown) => void };
const pending = new Map<string, Resolver[]>();
type Pending = { path: string; ext: string; resolvers: Resolver[] };
const pending = new Map<string, Pending>();
const cache = new Map<string, string>();
let flushTimer: ReturnType<typeof setTimeout> | null = null;
const FLUSH_DELAY_MS = 30;
const MAX_BATCH = 200;
function normExt(ext: string | undefined): string {
return (ext || 'txt').replace(/^\.+/, '').trim() || 'txt';
}
function keyFor(path: string, ext: string): string {
return `${ext}\n${path}`;
}
function scheduleFlush() {
if (flushTimer) return;
flushTimer = setTimeout(flush, FLUSH_DELAY_MS);
@@ -20,55 +31,69 @@ async function flush() {
flushTimer = null;
if (pending.size === 0) return;
// Drain up to MAX_BATCH paths; if more arrived, reschedule.
const paths: string[] = [];
for (const path of pending.keys()) {
paths.push(path);
if (paths.length >= MAX_BATCH) break;
// Drain up to MAX_BATCH entries; if more arrived, reschedule.
const keys: string[] = [];
for (const key of pending.keys()) {
keys.push(key);
if (keys.length >= MAX_BATCH) break;
}
const batchResolvers = paths.map(p => ({ path: p, resolvers: pending.get(p)! }));
for (const p of paths) pending.delete(p);
const drained = keys.map(k => pending.get(k)!);
for (const k of keys) pending.delete(k);
try {
const res = await apiClient.post('/api/caption/getBatch', { imgPaths: paths });
const captions: Record<string, string> = res.data?.captions ?? {};
for (const { path, resolvers } of batchResolvers) {
const value = captions[path] ?? '';
cache.set(path, value);
for (const r of resolvers) r.resolve(value);
}
} catch (err) {
for (const { resolvers } of batchResolvers) {
for (const r of resolvers) r.reject(err);
}
// Group by extension; each extension is a separate batch request.
const byExt = new Map<string, Pending[]>();
for (const entry of drained) {
const group = byExt.get(entry.ext);
if (group) group.push(entry);
else byExt.set(entry.ext, [entry]);
}
await Promise.all(
Array.from(byExt.entries()).map(async ([ext, entries]) => {
const paths = entries.map(e => e.path);
try {
const res = await apiClient.post('/api/caption/getBatch', { imgPaths: paths, ext });
const captions: Record<string, string> = res.data?.captions ?? {};
for (const { path, ext: e, resolvers } of entries) {
const value = captions[path] ?? '';
cache.set(keyFor(path, e), value);
for (const r of resolvers) r.resolve(value);
}
} catch (err) {
for (const { resolvers } of entries) {
for (const r of resolvers) r.reject(err);
}
}
}),
);
if (pending.size > 0) scheduleFlush();
}
function requestCaption(path: string, signal?: AbortSignal): Promise<string> {
function requestCaption(path: string, ext: string, signal?: AbortSignal): Promise<string> {
return new Promise((resolve, reject) => {
if (signal?.aborted) {
reject(new DOMException('Aborted', 'AbortError'));
return;
}
const key = keyFor(path, ext);
const resolver: Resolver = { resolve, reject };
const list = pending.get(path);
if (list) {
list.push(resolver);
const entry = pending.get(key);
if (entry) {
entry.resolvers.push(resolver);
} else {
pending.set(path, [resolver]);
pending.set(key, { path, ext, resolvers: [resolver] });
}
if (signal) {
const onAbort = () => {
// Remove this resolver from the pending batch. If no other card is
// still waiting on the same path, drop the path entirely so the next
// still waiting on the same key, drop the entry entirely so the next
// batch doesn't include it.
const arr = pending.get(path);
if (arr) {
const idx = arr.indexOf(resolver);
if (idx >= 0) arr.splice(idx, 1);
if (arr.length === 0) pending.delete(path);
const e = pending.get(key);
if (e) {
const idx = e.resolvers.indexOf(resolver);
if (idx >= 0) e.resolvers.splice(idx, 1);
if (e.resolvers.length === 0) pending.delete(key);
}
reject(new DOMException('Aborted', 'AbortError'));
};
@@ -78,19 +103,20 @@ function requestCaption(path: string, signal?: AbortSignal): Promise<string> {
});
}
export function invalidateCaption(path: string) {
cache.delete(path);
export function invalidateCaption(path: string, ext?: string) {
cache.delete(keyFor(path, normExt(ext)));
}
export function setCachedCaption(path: string, caption: string) {
cache.set(path, caption);
export function setCachedCaption(path: string, caption: string, ext?: string) {
cache.set(keyFor(path, normExt(ext)), caption);
}
// Fetches caption for a path, using the module-level batcher + cache.
// `refreshKey` busts the cache (e.g. after external edits or auto-captioning poll).
export default function useCaptionBatch(imgPath: string | null, refreshKey: number = 0) {
const [caption, setCaption] = useState<string>(() => (imgPath ? (cache.get(imgPath) ?? '') : ''));
const [isLoaded, setIsLoaded] = useState<boolean>(() => Boolean(imgPath && cache.has(imgPath)));
export default function useCaptionBatch(imgPath: string | null, refreshKey: number = 0, ext: string = 'txt') {
const captionExt = normExt(ext);
const [caption, setCaption] = useState<string>(() => (imgPath ? (cache.get(keyFor(imgPath, captionExt)) ?? '') : ''));
const [isLoaded, setIsLoaded] = useState<boolean>(() => Boolean(imgPath && cache.has(keyFor(imgPath, captionExt))));
const lastPathRef = useRef<string | null>(null);
useEffect(() => {
@@ -100,9 +126,9 @@ export default function useCaptionBatch(imgPath: string | null, refreshKey: numb
return;
}
if (refreshKey > 0) invalidateCaption(imgPath);
if (refreshKey > 0) invalidateCaption(imgPath, captionExt);
const cached = cache.get(imgPath);
const cached = cache.get(keyFor(imgPath, captionExt));
if (cached !== undefined) {
setCaption(cached);
setIsLoaded(true);
@@ -114,7 +140,7 @@ export default function useCaptionBatch(imgPath: string | null, refreshKey: numb
const controller = new AbortController();
lastPathRef.current = imgPath;
setIsLoaded(false);
requestCaption(imgPath, controller.signal)
requestCaption(imgPath, captionExt, controller.signal)
.then(value => {
if (cancelled || lastPathRef.current !== imgPath) return;
setCaption(value);
@@ -130,7 +156,7 @@ export default function useCaptionBatch(imgPath: string | null, refreshKey: numb
cancelled = true;
controller.abort();
};
}, [imgPath, refreshKey]);
}, [imgPath, refreshKey, captionExt]);
return { caption, isLoaded, setCaption };
}

View File

@@ -272,6 +272,7 @@ export interface CaptionProcessConfig {
max_res?: number;
max_new_tokens?: number;
fixed_caption?: string;
caption_extension?: string;
}
}