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:
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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}
|
||||
/>
|
||||
</>
|
||||
);
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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));
|
||||
};
|
||||
}, []);
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -21,6 +21,7 @@ export const defaultCaptionJobConfig: CaptionJobConfig = {
|
||||
extensions: ['mp3', 'wav', 'flac', 'ogg'],
|
||||
path_to_caption: '',
|
||||
recaption: false,
|
||||
caption_extension: 'txt',
|
||||
},
|
||||
},
|
||||
],
|
||||
|
||||
21
ui/src/helpers/datasetTemplates.ts
Normal file
21
ui/src/helpers/datasetTemplates.ts
Normal 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;
|
||||
@@ -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 };
|
||||
}
|
||||
|
||||
@@ -272,6 +272,7 @@ export interface CaptionProcessConfig {
|
||||
max_res?: number;
|
||||
max_new_tokens?: number;
|
||||
fixed_caption?: string;
|
||||
caption_extension?: string;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user