variable changes and some cleanup

This commit is contained in:
Manan17 2026-03-03 09:34:32 +00:00
commit f04c684d8a
19 changed files with 93 additions and 111 deletions

View file

@ -30,11 +30,6 @@ from trl import SFTTrainer, SFTConfig
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# Process-level flag: set True after CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec).
# Once CUDA has been used for audio processing, fork-based multiprocessing (num_proc>1)
# deadlocks because forked children inherit CUDA's internal thread locks.
# This flag is never reset — once contaminated, the process stays contaminated.
_CUDA_AUDIO_PREPROCESSING_DONE = False
@dataclass
class TrainingProgress:
@ -75,6 +70,7 @@ class UnslothTrainer:
self.is_audio = False
self.is_audio_vlm = False # Multimodal model (e.g. Gemma 3N) trained on audio data
self._audio_type = None # 'csm', 'whisper', 'snac', 'xcodec2', 'bicodec', 'dac'
self._cuda_audio_used = False # Set once after audio CUDA preprocessing; never cleared
self._spark_tts_repo_dir = None # Path to downloaded Spark-TTS repo (for BiCodecTokenizer)
self.model_name = None
@ -335,7 +331,7 @@ class UnslothTrainer:
max_seq_length: int = 2048,
load_in_4bit: bool = True,
hf_token: Optional[str] = None,
is_dataset_multimodal: bool = False,
is_dataset_image: bool = False,
is_dataset_audio: bool = False) -> bool:
"""Load model for training (supports both text and vision models)"""
try:
@ -381,14 +377,14 @@ class UnslothTrainer:
# Uses FastModel + SFTTrainer with audio collator
self.is_audio_vlm = not self.is_audio and is_vision_model(model_name) and is_dataset_audio
# VLM: vision model with image dataset (mutually exclusive with audio VLM)
self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_multimodal
self.is_vlm = not self.is_audio and not self.is_audio_vlm and is_vision_model(model_name) and is_dataset_image
self.model_name = model_name
self.max_seq_length = max_seq_length
logger.info(f"Audio type: {self._audio_type}")
if not self.is_audio:
logger.info(f"Model architecture is vision: {is_vision_model(model_name)}")
logger.info(f"Dataset is multimodal: {is_dataset_multimodal}, audio: {is_dataset_audio}")
logger.info(f"Dataset has images: {is_dataset_image}, audio: {is_dataset_audio}")
logger.info(f"Using VLM path: {self.is_vlm}, Audio VLM: {self.is_audio_vlm}")
# Reset training state for new run
@ -555,10 +551,26 @@ class UnslothTrainer:
print("Model loaded successfully")
return True
except OSError as e:
if "could not get source code" in str(e) and not getattr(self, '_source_code_retried', False):
# Unsloth's patching can leave stale state that makes
# inspect.getsource() fail when switching model families
# (e.g. gemma3 → gemma3n). The load always succeeds on a
# second attempt because the failed first call's partial
# imports clean up the stale state as a side effect.
self._source_code_retried = True
print(f"\n'could not get source code' — retrying once...\n")
return self.load_model(model_name, max_seq_length, load_in_4bit, hf_token,
is_dataset_image, is_dataset_audio)
logger.error(f"Error loading model: {e}")
self._update_progress(error=str(e), is_training=False)
return False
except Exception as e:
logger.error(f"Error loading model: {e}")
self._update_progress(error=str(e), is_training=False)
return False
finally:
self._source_code_retried = False
def prepare_model_for_training(self,
use_lora: bool = True,
@ -1204,9 +1216,7 @@ class UnslothTrainer:
import gc
gc.collect()
torch.cuda.empty_cache()
global _CUDA_AUDIO_PREPROCESSING_DONE
_CUDA_AUDIO_PREPROCESSING_DONE = True
self._cuda_audio_used = True
if not processed_examples:
raise ValueError(
@ -1394,9 +1404,7 @@ class UnslothTrainer:
import gc
gc.collect()
torch.cuda.empty_cache()
global _CUDA_AUDIO_PREPROCESSING_DONE
_CUDA_AUDIO_PREPROCESSING_DONE = True
self._cuda_audio_used = True
if not processed_examples:
raise ValueError(
@ -1579,13 +1587,7 @@ class UnslothTrainer:
import gc
gc.collect()
torch.cuda.empty_cache()
# Mark process as CUDA-contaminated from audio preprocessing.
# Fork-based multiprocessing (num_proc>1) will deadlock after this
# because forked children inherit CUDA's internal thread locks from
# Whisper/DAC processing that can't be released.
global _CUDA_AUDIO_PREPROCESSING_DONE
_CUDA_AUDIO_PREPROCESSING_DONE = True
self._cuda_audio_used = True
if not processed_examples:
raise ValueError(
@ -2236,7 +2238,7 @@ class UnslothTrainer:
"output_dir": output_dir,
"report_to": ["wandb"] if training_args.get('enable_wandb', False) else "none",
"include_num_input_tokens_seen": True, # Enable token counting
"dataset_num_proc": safe_num_proc(max(1, os.cpu_count() // 4)),
"dataset_num_proc": 1 if (self.is_audio or self.is_audio_vlm or self._cuda_audio_used) else safe_num_proc(max(1, os.cpu_count() // 4)),
"max_seq_length": training_args.get('max_seq_length', 2048),
}
@ -2424,19 +2426,11 @@ class UnslothTrainer:
try:
from unsloth.chat_templates import train_on_responses_only
# After CUDA-heavy audio preprocessing (Whisper/DAC/BiCodec/SNAC),
# fork-based multiprocessing deadlocks because children inherit
# CUDA's internal thread locks. Use single-process mode instead.
toro_num_proc = config_args.get("dataset_num_proc", safe_num_proc(max(1, os.cpu_count() // 4)))
if _CUDA_AUDIO_PREPROCESSING_DONE:
toro_num_proc = 1
print("Using single-process train_on_responses (CUDA audio preprocessing detected)\n")
self.trainer = train_on_responses_only(
self.trainer,
instruction_part=instruction_part,
response_part=response_part,
num_proc=toro_num_proc,
num_proc=config_args["dataset_num_proc"],
)
print("Train on responses only configured successfully\n")

View file

@ -116,7 +116,7 @@ class TrainingBackend:
train_split: str = "train",
eval_split: str = None,
eval_steps: float = 0.00,
is_dataset_multimodal: bool = False,
is_dataset_image: bool = False,
is_dataset_audio: bool = False) -> bool:
"""
Start training.
@ -146,9 +146,6 @@ class TrainingBackend:
import torch as _torch
if _torch.cuda.is_available():
_torch.cuda.synchronize()
# Reset torch dynamo/compiler caches — Unsloth's compiled SFTTrainer
# and model.for_training() set class-level state that persists between
# runs (e.g. BiCodec text trainer pollutes subsequent VLM runs).
_torch._dynamo.reset()
_torch.compiler.reset()
import gc
@ -182,7 +179,7 @@ class TrainingBackend:
max_seq_length=max_seq_length,
load_in_4bit=load_in_4bit if use_lora_actual else False, # Only 4bit for LoRA
hf_token=hf_token if hf_token.strip() else None,
is_dataset_multimodal=is_dataset_multimodal,
is_dataset_image=is_dataset_image,
is_dataset_audio=is_dataset_audio,
)

View file

@ -27,7 +27,7 @@ class CheckFormatResponse(BaseModel):
requires_manual_mapping: bool
detected_format: str
columns: List[str]
is_multimodal: bool = False
is_image: bool = False
is_audio: bool = False
multimodal_columns: Optional[List[str]] = None
suggested_mapping: Optional[Dict[str, str]] = None

View file

@ -65,7 +65,7 @@ class TrainingStartRequest(BaseModel):
finetune_language_layers: bool = Field(False, description="Finetune language layers")
finetune_attention_modules: bool = Field(False, description="Finetune attention modules")
finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules")
is_dataset_multimodal: bool = Field(False, description="Whether the dataset contains multimodal (image) data")
is_dataset_image: bool = Field(False, description="Whether the dataset contains image data")
is_dataset_audio: bool = Field(False, description="Whether the dataset contains audio data")
# Logging parameters

View file

@ -183,7 +183,7 @@ def check_format(request: CheckFormatRequest):
# Run lightweight format check on the preview slice
result = check_dataset_format(preview_slice, is_vlm=request.is_vlm)
logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_multimodal={result.get('is_multimodal', False)}")
logger.info(f"Format check result: requires_mapping={result['requires_manual_mapping']}, format={result['detected_format']}, is_image={result.get('is_image', False)}")
# Generate preview samples
preview_samples = None
@ -211,7 +211,7 @@ def check_format(request: CheckFormatRequest):
requires_manual_mapping=result["requires_manual_mapping"],
detected_format=result["detected_format"],
columns=result["columns"],
is_multimodal=result.get("is_multimodal", False),
is_image=result.get("is_image", False),
is_audio=result.get("is_audio", False),
multimodal_columns=result.get("multimodal_columns"),
suggested_mapping=result.get("suggested_mapping"),

View file

@ -178,7 +178,7 @@ async def start_training(
"finetune_language_layers": request.finetune_language_layers,
"finetune_attention_modules": request.finetune_attention_modules,
"finetune_mlp_modules": request.finetune_mlp_modules,
"is_dataset_multimodal": request.is_dataset_multimodal,
"is_dataset_image": request.is_dataset_image,
"is_dataset_audio": request.is_dataset_audio,
"enable_wandb": request.enable_wandb,
"wandb_token": request.wandb_token or "",

View file

@ -67,8 +67,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
multimodal_info = detect_multimodal_dataset(dataset)
is_audio = multimodal_info.get("is_audio", False)
if multimodal_info["is_multimodal"] and not is_audio:
is_vlm = True # Route to VLM detection for image datasets only
if multimodal_info["is_image"]:
is_vlm = True # Route to VLM detection for image datasets
# Common audio fields for all return paths
audio_fields = {
@ -88,7 +88,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None,
"detected_image_column": vlm_structure.get("image_column"),
"detected_text_column": vlm_structure.get("text_column"),
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_columns": multimodal_info.get("multimodal_columns"),
**audio_fields,
}
@ -105,7 +105,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None,
"detected_image_column": None,
"detected_text_column": multimodal_info.get("detected_text_column"),
"is_multimodal": True,
"is_image": False,
"multimodal_columns": multimodal_info.get("audio_columns"),
**audio_fields,
}
@ -124,7 +124,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": heuristic_mapping,
"detected_image_column": None,
"detected_text_column": None,
"is_multimodal": False,
"is_image": False,
"multimodal_columns": None,
**audio_fields,
}
@ -136,7 +136,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None,
"detected_image_column": None,
"detected_text_column": None,
"is_multimodal": False,
"is_image": False,
"multimodal_columns": None,
**audio_fields,
}
@ -149,7 +149,7 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None,
"detected_image_column": None,
"detected_text_column": None,
"is_multimodal": False,
"is_image": False,
"multimodal_columns": None,
**audio_fields,
}
@ -278,7 +278,7 @@ def format_dataset(
"chat_column": chat_column,
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": [f"Applied user-provided column mapping ({format_type}): {custom_format_mapping}"]
}
@ -290,7 +290,7 @@ def format_dataset(
"chat_column": None,
"is_standardized": False,
"requires_manual_mapping": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": [f"Failed to apply user mapping: {e}"]
}
@ -301,7 +301,7 @@ def format_dataset(
warnings = []
# Add multimodal warning if detected
if multimodal_info["is_multimodal"]:
if multimodal_info["is_image"]:
warnings.append(
f"Multimodal dataset detected. Found columns: {multimodal_info['multimodal_columns']}"
)
@ -318,7 +318,7 @@ def format_dataset(
"chat_column": None,
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -338,7 +338,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -351,7 +351,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": False,
"requires_manual_mapping": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -364,7 +364,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -415,7 +415,7 @@ def format_dataset(
"chat_column": "conversations",
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -438,7 +438,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -453,7 +453,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": False,
"requires_manual_mapping": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -469,7 +469,7 @@ def format_dataset(
"chat_column": None,
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -492,7 +492,7 @@ def format_dataset(
"chat_column": None,
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -506,7 +506,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": False,
"requires_manual_mapping": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -523,7 +523,7 @@ def format_dataset(
"chat_column": "conversations",
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -541,7 +541,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -554,7 +554,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": []
}
@ -575,7 +575,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": True,
"requires_manual_mapping": False,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -589,7 +589,7 @@ def format_dataset(
"chat_column": detected["chat_column"],
"is_standardized": False,
"requires_manual_mapping": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"warnings": warnings
}
@ -675,7 +675,7 @@ def format_and_template_dataset(
"final_format": "vlm_messages",
"chat_column": "messages",
"is_vlm": True,
"is_multimodal": True,
"is_image": True,
"multimodal_info": multimodal_info,
"success": True,
"requires_manual_mapping": False,
@ -797,7 +797,7 @@ def format_and_template_dataset(
"final_format": "vlm_messages",
"chat_column": "messages",
"is_vlm": True,
"is_multimodal": multimodal_info["is_multimodal"],
"is_image": multimodal_info["is_image"],
"multimodal_info": multimodal_info,
"vlm_structure": vlm_structure,
"success": True,
@ -826,7 +826,7 @@ def format_and_template_dataset(
# Gemma emits a leading <bos> that must be stripped for text-only chatml/sharegpt.
is_alpaca = format_type == "alpaca" or (format_type == "auto" and dataset_info["detected_format"] == "alpaca")
is_gemma = "gemma" in model_name.lower()
if is_gemma and not dataset_info["is_multimodal"] and not is_alpaca:
if is_gemma and not dataset_info["is_image"] and not is_alpaca:
remove_bos_prefix = True
template_result = apply_chat_template_to_dataset(
dataset_info=dataset_info,

View file

@ -334,7 +334,7 @@ def detect_multimodal_dataset(dataset):
Returns:
dict: {
"is_multimodal": bool,
"is_image": bool,
"multimodal_columns": list of column names containing image data,
"modality_types": list of detected types (e.g., ["image", "audio"]),
"is_audio": bool,
@ -427,7 +427,7 @@ def detect_multimodal_dataset(dataset):
break
return {
"is_multimodal": len(multimodal_columns) > 0 or is_audio,
"is_image": len(multimodal_columns) > 0,
"multimodal_columns": multimodal_columns,
"modality_types": list(modality_types),
"is_audio": is_audio,

View file

@ -423,14 +423,13 @@ def safe_num_proc(desired: Optional[int] = None) -> int:
"""
Return a safe ``num_proc`` for ``dataset.map()`` calls.
Fork-based multiprocessing deadlocks when CUDA has already been
initialized (e.g. after inference). This helper detects that case
and forces ``num_proc=1``.
On multi-GPU machines the NVIDIA driver spawns extra background threads,
making ``os.fork()`` prone to deadlocks when many workers are created.
This helper caps ``num_proc`` to 4 on such machines.
On single-GPU (or CPU-only) machines the original value is returned
unchanged.
Args:
desired: The num_proc you *want*. If None, auto-computes from
``os.cpu_count()``.
@ -443,14 +442,6 @@ def safe_num_proc(desired: Optional[int] = None) -> int:
if desired is None or not isinstance(desired, int):
desired = max(1, os.cpu_count() // 3)
# After inference, CUDA is initialized — forking will deadlock.
try:
import torch
if torch.cuda.is_initialized():
return 1
except ImportError:
pass
if get_physical_gpu_count() > 1:
capped = min(4, desired)
print(

View file

@ -62,10 +62,10 @@ export function DatasetPreviewDialog({
);
const { isStarting, startError, startTrainingRun } = useTrainingActions();
// If the backend reports multimodal data, treat as VLM even if the prop
// hasn't caught up yet (isDatasetMultimodal may still be null in the store).
// If the backend reports image data, treat as VLM even if the prop
// hasn't caught up yet (isDatasetImage may still be null in the store).
const effectiveIsAudio = !!data?.is_audio;
const effectiveIsVlm = !effectiveIsAudio && (isVlm || !!data?.is_multimodal);
const effectiveIsVlm = isVlm || !!data?.is_image;
const hasHeuristicMapping = !data?.requires_manual_mapping && !!data?.suggested_mapping;
const mappingEnabled = !!data?.requires_manual_mapping || hasHeuristicMapping;

View file

@ -114,7 +114,7 @@ function SliderRow({
export function ParamsSection(): ReactElement {
const store = useTrainingConfigStore();
const isLora = store.trainingMethod !== "full";
const showVisionLora = store.isVisionModel && store.isDatasetMultimodal === true;
const showVisionLora = store.isVisionModel && store.isDatasetImage === true;
const [loraOpen, setLoraOpen] = useState(false);
const [hyperOpen, setHyperOpen] = useState(false);
const maxStepsSliderMax = Math.max(500, store.maxSteps, 30);

View file

@ -42,7 +42,7 @@ export function TrainingSection() {
const store = useTrainingConfigStore();
const { isStarting, startError, startTrainingRun } = useTrainingActions();
const isIncompatible =
!store.isVisionModel && !store.isDatasetAudio && store.isDatasetMultimodal === true;
!store.isVisionModel && store.isDatasetImage === true;
const fileInputRef = useRef<HTMLInputElement>(null);
const handleFileUpload = (e: React.ChangeEvent<HTMLInputElement>) => {

View file

@ -93,7 +93,7 @@ export function StudioPage(): ReactElement {
datasetSplit={config.datasetSplit}
mode={dialogMode}
initialData={dialogInitial}
isVlm={config.isVisionModel && config.isDatasetMultimodal === true}
isVlm={config.isVisionModel && config.isDatasetImage === true}
/>
{canGoBack && (

View file

@ -57,7 +57,7 @@ export function buildTrainingStartPayload(
finetune_language_layers: config.finetuneLanguageLayers,
finetune_attention_modules: config.finetuneAttentionModules,
finetune_mlp_modules: config.finetuneMLPModules,
is_dataset_multimodal: !!config.isDatasetMultimodal,
is_dataset_image: !!config.isDatasetImage,
is_dataset_audio: config.isDatasetAudio,
enable_wandb: config.enableWandb,
wandb_token: config.enableWandb ? config.wandbToken.trim() || null : null,

View file

@ -36,7 +36,7 @@ export function useTrainingActions() {
try {
const datasetName = getDatasetName(config);
let isVlm = config.isVisionModel && config.isDatasetMultimodal === true;
let isVlm = config.isVisionModel && config.isDatasetImage === true;
if (datasetName) {
const check = await checkDatasetFormat({
@ -47,17 +47,17 @@ export function useTrainingActions() {
isVlm,
});
// Backend auto-detects multimodal/audio from dataset content.
// Backend auto-detects image/audio from dataset content.
// Sync these flags into the store so buildTrainingStartPayload picks them up.
const isAudio = !!check.is_audio;
const isMultimodal = !!check.is_multimodal;
const isImage = !!check.is_image;
if (isMultimodal && config.isVisionModel) {
if (isImage && config.isVisionModel) {
isVlm = true;
}
if (isMultimodal !== config.isDatasetMultimodal || isAudio !== config.isDatasetAudio) {
if (isImage !== config.isDatasetImage || isAudio !== config.isDatasetAudio) {
useTrainingConfigStore.setState({
isDatasetMultimodal: isMultimodal,
isDatasetImage: isImage,
isDatasetAudio: isAudio,
});
}

View file

@ -35,7 +35,7 @@ const initialState: TrainingConfigState = {
modelDefaultsError: null,
modelDefaultsAppliedFor: null,
isCheckingDataset: false,
isDatasetMultimodal: null,
isDatasetImage: null,
isDatasetAudio: false,
...DEFAULT_HYPERPARAMS,
};
@ -57,7 +57,7 @@ const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set
"modelDefaultsError",
"modelDefaultsAppliedFor",
"isCheckingDataset",
"isDatasetMultimodal",
"isDatasetImage",
"isDatasetAudio",
"trainOnCompletions",
]);
@ -116,9 +116,9 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
_trainOnCompletionsManuallySet = false;
const patch = mapBackendModelConfigToTrainingPatch(modelDetails.config);
// If vision model + multimodal dataset already known, override
// If vision model + image dataset already known, override
// trainOnCompletions to false regardless of backend default.
if (modelDetails.is_vision && get().isDatasetMultimodal === true) {
if (modelDetails.is_vision && get().isDatasetImage === true) {
patch.trainOnCompletions = false;
}
@ -174,16 +174,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
})
.then((res) => {
if (controller.signal.aborted) return;
const isMultimodal = !!res.is_multimodal;
const isImage = !!res.is_image;
const isAudio = !!res.is_audio;
const updates: Record<string, unknown> = {
isDatasetMultimodal: isMultimodal,
isDatasetImage: isImage,
isDatasetAudio: isAudio,
isCheckingDataset: false,
};
if (!_trainOnCompletionsManuallySet) {
const { isVisionModel } = get();
if (isVisionModel && isMultimodal) {
if (isVisionModel && isImage) {
updates.trainOnCompletions = false;
}
}
@ -191,7 +191,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
})
.catch(() => {
if (controller.signal.aborted) return;
set({ isDatasetMultimodal: null, isCheckingDataset: false });
set({ isDatasetImage: null, isCheckingDataset: false });
});
};
@ -261,7 +261,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
datasetSplit: null,
datasetEvalSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isDatasetImage: null,
isCheckingDataset: false,
});
},
@ -274,7 +274,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
datasetSplit: null,
datasetEvalSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isDatasetImage: null,
isCheckingDataset: false,
});
},
@ -282,7 +282,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
set({
datasetSplit,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isDatasetImage: null,
isCheckingDataset: false,
});
@ -298,7 +298,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
ensureDatasetChecked: () => {
const state = get();
if (state.isCheckingDataset) return;
if (state.isDatasetMultimodal !== null) return;
if (state.isDatasetImage !== null) return;
const datasetName =
state.datasetSource === "huggingface"

View file

@ -38,7 +38,7 @@ export interface TrainingStartRequest {
finetune_language_layers: boolean;
finetune_attention_modules: boolean;
finetune_mlp_modules: boolean;
is_dataset_multimodal: boolean;
is_dataset_image: boolean;
is_dataset_audio: boolean;
enable_wandb: boolean;
wandb_token: string | null;

View file

@ -59,7 +59,7 @@ export interface TrainingConfigState {
modelDefaultsError: string | null;
modelDefaultsAppliedFor: string | null;
isCheckingDataset: boolean;
isDatasetMultimodal: boolean | null;
isDatasetImage: boolean | null;
isDatasetAudio: boolean;
finetuneVisionLayers: boolean;
finetuneLanguageLayers: boolean;

View file

@ -9,7 +9,7 @@ export type CheckFormatResponse = {
detected_speaker_column?: string | null;
preview_samples?: Record<string, unknown>[] | null;
total_rows?: number | null;
is_multimodal?: boolean;
is_image?: boolean;
is_audio?: boolean;
multimodal_columns?: string[] | null;
};