diff --git a/extensions_built_in/captioner/AceStepCaptioner.py b/extensions_built_in/captioner/AceStepCaptioner.py
index 8ecb1ff8..75897a40 100644
--- a/extensions_built_in/captioner/AceStepCaptioner.py
+++ b/extensions_built_in/captioner/AceStepCaptioner.py
@@ -1,3 +1,5 @@
+from typing import Optional
+
import librosa
import numpy as np
import torch
@@ -9,7 +11,7 @@ from optimum.quanto import freeze
from toolkit.basic import flush
from toolkit.util.quantize import quantize, get_qtype
-from .BaseCaptioner import BaseCaptioner
+from .BaseCaptioner import BaseCaptioner, CaptionConfig
import transformers
import logging
import warnings
@@ -95,7 +97,16 @@ def analyze_audio(audio_path):
}
+class AceStepCaptionConfig(CaptionConfig):
+ def __init__(self, **kwargs):
+ super().__init__(**kwargs)
+ self.fixed_caption: Optional[str] = kwargs.get("fixed_caption", None)
+
+
class AceStepCaptioner(BaseCaptioner):
+ caption_config_class = AceStepCaptionConfig
+ caption_config: AceStepCaptionConfig
+
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super(AceStepCaptioner, self).__init__(process_id, job, config, **kwargs)
@@ -118,27 +129,31 @@ class AceStepCaptioner(BaseCaptioner):
)
if self.caption_config.low_vram:
self.model.to("cpu")
+
+ self.model2 = None
+ self.processor2 = None
- # load captioner model
- self.print_and_status_update("Loading captioner model")
- self.model2 = Qwen2_5OmniForConditionalGeneration.from_pretrained(
- self.caption_config.model_name_or_path2,
- dtype=self.torch_dtype,
- device_map="cpu",
- )
- self.model2.to(self.device_torch)
- self.model2.disable_talker()
- if self.caption_config.quantize:
- self.print_and_status_update("Quantizing captioner model")
- quantize(self.model2, weights=get_qtype(self.caption_config.qtype))
- freeze(self.model2)
- flush()
- self.processor2 = Qwen2_5OmniProcessor.from_pretrained(
- self.caption_config.model_name_or_path2,
- )
+ if self.caption_config.fixed_caption is not None:
+ # load captioner model
+ self.print_and_status_update("Loading captioner model")
+ self.model2 = Qwen2_5OmniForConditionalGeneration.from_pretrained(
+ self.caption_config.model_name_or_path2,
+ dtype=self.torch_dtype,
+ device_map="cpu",
+ )
+ self.model2.to(self.device_torch)
+ self.model2.disable_talker()
+ if self.caption_config.quantize:
+ self.print_and_status_update("Quantizing captioner model")
+ quantize(self.model2, weights=get_qtype(self.caption_config.qtype))
+ freeze(self.model2)
+ flush()
+ self.processor2 = Qwen2_5OmniProcessor.from_pretrained(
+ self.caption_config.model_name_or_path2,
+ )
- if self.caption_config.low_vram:
- self.model2.to("cpu")
+ if self.caption_config.low_vram:
+ self.model2.to("cpu")
flush()
def run_qwen_audio(self, model, processor, audio_data, sr, prompt_text):
@@ -228,7 +243,10 @@ class AceStepCaptioner(BaseCaptioner):
lyrics = lyrics.split("# Lyrics")[1].strip()
# get the caption from the audio
- caption = self.get_audio_caption(audio_data)
+ if self.caption_config.fixed_caption is not None:
+ caption = self.caption_config.fixed_caption
+ else:
+ caption = self.get_audio_caption(audio_data)
output = f"
\n{caption}\n\n"
output += f"\n{lyrics}\n\n"
diff --git a/extensions_built_in/captioner/BaseCaptioner.py b/extensions_built_in/captioner/BaseCaptioner.py
index 3be8a522..bd946292 100644
--- a/extensions_built_in/captioner/BaseCaptioner.py
+++ b/extensions_built_in/captioner/BaseCaptioner.py
@@ -46,6 +46,8 @@ class CaptionConfig:
class BaseCaptioner(BaseExtensionProcess):
+ caption_config_class = CaptionConfig
+
def __init__(self, process_id: int, job, config: OrderedDict, **kwargs):
super(BaseCaptioner, self).__init__(process_id, job, config, **kwargs)
self.sqlite_db_path = self.config.get("sqlite_db_path", "./aitk_db.db")
@@ -74,7 +76,7 @@ class BaseCaptioner(BaseExtensionProcess):
self._stop_watcher_started = False
# self.start_stop_watcher(interval_sec=2.0)
- self.caption_config = CaptionConfig(**self.get_conf("caption", {}))
+ self.caption_config = self.caption_config_class(**self.get_conf("caption", {}))
self.model = None
self.processor = None
self.model2 = None
@@ -85,19 +87,20 @@ class BaseCaptioner(BaseExtensionProcess):
def run(self):
super(BaseCaptioner, self).run()
- self.start_stop_watcher()
- self.update_status("running", "Loading Model")
- self.load_model()
- self.update_status("running", "Looking for files")
- self.find_files()
- self.update_status("running", f"Captioning {len(self.file_paths)} files")
- self.run_caption_loop()
- self.update_status("completed", "Captioning completed")
- print("")
+ with torch.no_grad():
+ self.start_stop_watcher()
+ self.update_status("running", "Loading Model")
+ self.load_model()
+ self.update_status("running", "Looking for files")
+ self.find_files()
+ self.update_status("running", f"Captioning {len(self.file_paths)} files")
+ self.run_caption_loop()
+ self.update_status("completed", "Captioning completed")
+ print("")
- print("****************************************************")
- print("Captioning complete")
- print("****************************************************")
+ print("****************************************************")
+ print("Captioning complete")
+ print("****************************************************")
def run_caption_loop(self):
for file_path in tqdm.tqdm(
diff --git a/ui/src/app/globals.css b/ui/src/app/globals.css
index 1ab7f24d..af1bdc97 100644
--- a/ui/src/app/globals.css
+++ b/ui/src/app/globals.css
@@ -46,7 +46,7 @@ body {
@layer components {
/* control */
.aitk-react-select-container .aitk-react-select__control {
- @apply flex w-full h-8 min-h-0 px-0 text-sm bg-gray-800 border border-gray-700 rounded-sm hover:border-gray-600 items-center;
+ @apply flex w-full h-8 min-h-0 px-0 text-sm bg-gray-950 dark:bg-gray-800 border border-gray-700 rounded-sm hover:border-gray-600 items-center;
}
/* selected label */
@@ -66,12 +66,12 @@ body {
/* menu */
.aitk-react-select-container .aitk-react-select__menu {
- @apply bg-gray-800 border border-gray-700;
+ @apply bg-gray-950 dark:bg-gray-800 border border-gray-700;
}
/* options */
.aitk-react-select-container .aitk-react-select__option {
- @apply text-sm text-gray-200 bg-gray-800 hover:bg-gray-700;
+ @apply text-sm text-gray-200 bg-gray-950 dark:bg-gray-800 hover:bg-gray-700;
}
/* indicator separator */
diff --git a/ui/src/components/CaptionDatasetModal.tsx b/ui/src/components/CaptionDatasetModal.tsx
index 2701eaa8..ebc5340a 100644
--- a/ui/src/components/CaptionDatasetModal.tsx
+++ b/ui/src/components/CaptionDatasetModal.tsx
@@ -111,6 +111,8 @@ export const CaptionDatasetModal: React.FC = () => {
});
};
+ const additionalSections = selectedCaptionOption?.additionalSections || [];
+
return (
@@ -159,7 +161,7 @@ export const CaptionDatasetModal: React.FC = () => {
required
/>
- {selectedCaptionOption?.additionalSections?.includes('caption.model_name_or_path2') && (
+ {additionalSections.includes('caption.model_name_or_path2') && (
{
/>
)}
+ {additionalSections.includes('caption.fixed_caption') && (
+
+ {
+ if (value?.trim() === '') {
+ //@ts-ignore
+ value = undefined;
+ }
+ setJobConfig(value, 'config.process[0].caption.fixed_caption');
+ }}
+ placeholder="Enter fixed caption (if you want the same caption for all audio files)"
+ />
+
+ )}
{
}}
options={quantizationOptions}
/>
- {selectedCaptionOption?.additionalSections?.includes('caption.max_res') && (
+ {additionalSections.includes('caption.max_res') && (
{
/>
)}
- {selectedCaptionOption?.additionalSections?.includes('caption.max_new_tokens') && (
+ {additionalSections.includes('caption.max_new_tokens') && (
{
- {selectedCaptionOption?.additionalSections?.includes('caption.caption_prompt') && (
+ {additionalSections.includes('caption.caption_prompt') && (
= ({
>
{/* Modal panel */}
e.stopPropagation()}
>
{/* Modal header */}
diff --git a/ui/src/components/formInputs.tsx b/ui/src/components/formInputs.tsx
index 74d07782..c498fa0d 100644
--- a/ui/src/components/formInputs.tsx
+++ b/ui/src/components/formInputs.tsx
@@ -12,7 +12,7 @@ const Select = dynamic(() => import('react-select'), { ssr: false });
const labelClasses = 'block text-xs mb-1 mt-2 text-gray-300';
const inputClasses =
- 'w-full text-sm px-3 py-1 bg-gray-800 border border-gray-700 rounded-sm focus:ring-2 focus:ring-gray-600 focus:border-transparent';
+ 'w-full text-sm px-3 py-1 bg-gray-950 dark:bg-gray-800 border border-gray-700 rounded-sm text-gray-100 placeholder:text-gray-500 focus:ring-2 focus:ring-gray-600 focus:border-transparent';
export interface InputProps {
label?: string;
@@ -584,7 +584,7 @@ export const SliderInput: React.FC = props => {
{showValue && (
-
+
{Number.isFinite(value) ? value : ''}
)}
diff --git a/ui/src/helpers/captionOptions.ts b/ui/src/helpers/captionOptions.ts
index 35392276..31dd7ffd 100644
--- a/ui/src/helpers/captionOptions.ts
+++ b/ui/src/helpers/captionOptions.ts
@@ -1,7 +1,7 @@
import { GroupedSelectOption, SelectOption } from "@/types";
type CaptionGroup = 'image' | 'music';
-type AdditionalSections = 'caption.model_name_or_path2' | 'caption.caption_prompt' | 'caption.max_res' | 'caption.max_new_tokens';
+type AdditionalSections = 'caption.model_name_or_path2' | 'caption.caption_prompt' | 'caption.max_res' | 'caption.max_new_tokens' | 'caption.fixed_caption';
export interface CaptionOption {
name: string;
@@ -41,6 +41,7 @@ export const captionerTypes: CaptionOption[] = [
],
additionalSections: [
'caption.model_name_or_path2',
+ 'caption.fixed_caption',
],
},
{
diff --git a/ui/src/types.ts b/ui/src/types.ts
index 7ed788e2..087362ad 100644
--- a/ui/src/types.ts
+++ b/ui/src/types.ts
@@ -269,6 +269,7 @@ export interface CaptionProcessConfig {
caption_prompt?: string;
max_res?: number;
max_new_tokens?: number;
+ fixed_caption?: string;
}
}