Merge pull request #163 from unslothai/fix/add-vision-and-dataset-to-local
Fix/add vision and dataset to local
This commit is contained in:
commit
b35798c8ff
1 changed files with 35 additions and 8 deletions
|
|
@ -46,11 +46,9 @@ let _modelConfigController: AbortController | null = null;
|
||||||
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
|
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
|
||||||
"modelType",
|
"modelType",
|
||||||
"isCheckingVision",
|
"isCheckingVision",
|
||||||
"isVisionModel",
|
|
||||||
"isLoadingModelDefaults",
|
"isLoadingModelDefaults",
|
||||||
"modelDefaultsError",
|
"modelDefaultsError",
|
||||||
"isCheckingDataset",
|
"isCheckingDataset",
|
||||||
"isDatasetMultimodal",
|
|
||||||
]);
|
]);
|
||||||
|
|
||||||
function partializePersistedState(
|
function partializePersistedState(
|
||||||
|
|
@ -104,8 +102,21 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
||||||
if (controller.signal.aborted) return;
|
if (controller.signal.aborted) return;
|
||||||
if (get().selectedModel !== modelName) return;
|
if (get().selectedModel !== modelName) return;
|
||||||
|
|
||||||
|
const patch = mapBackendModelConfigToTrainingPatch(modelDetails.config);
|
||||||
|
|
||||||
|
// train_on_responses_only: true for LLMs, true for VLM+text dataset, false for VLM+vision dataset
|
||||||
|
if (!modelDetails.is_vision) {
|
||||||
|
patch.trainOnCompletions = true;
|
||||||
|
} else {
|
||||||
|
const datasetMultimodal = get().isDatasetMultimodal;
|
||||||
|
if (datasetMultimodal !== null) {
|
||||||
|
patch.trainOnCompletions = !datasetMultimodal;
|
||||||
|
}
|
||||||
|
// if dataset not yet checked, leave YAML default
|
||||||
|
}
|
||||||
|
|
||||||
set({
|
set({
|
||||||
...mapBackendModelConfigToTrainingPatch(modelDetails.config),
|
...patch,
|
||||||
isVisionModel: modelDetails.is_vision,
|
isVisionModel: modelDetails.is_vision,
|
||||||
isLoadingModelDefaults: false,
|
isLoadingModelDefaults: false,
|
||||||
isCheckingVision: false,
|
isCheckingVision: false,
|
||||||
|
|
@ -129,10 +140,20 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
||||||
void checkVisionModel(modelName)
|
void checkVisionModel(modelName)
|
||||||
.then((isVision) => {
|
.then((isVision) => {
|
||||||
if (get().selectedModel !== modelName) return;
|
if (get().selectedModel !== modelName) return;
|
||||||
set({
|
const updates: Record<string, unknown> = {
|
||||||
isVisionModel: isVision,
|
isVisionModel: isVision,
|
||||||
isCheckingVision: false,
|
isCheckingVision: false,
|
||||||
});
|
};
|
||||||
|
// train_on_responses_only default
|
||||||
|
if (!isVision) {
|
||||||
|
updates.trainOnCompletions = true;
|
||||||
|
} else {
|
||||||
|
const datasetMultimodal = get().isDatasetMultimodal;
|
||||||
|
if (datasetMultimodal !== null) {
|
||||||
|
updates.trainOnCompletions = !datasetMultimodal;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
set(updates as Partial<TrainingConfigState>);
|
||||||
})
|
})
|
||||||
.catch(() => {
|
.catch(() => {
|
||||||
if (get().selectedModel !== modelName) return;
|
if (get().selectedModel !== modelName) return;
|
||||||
|
|
@ -247,10 +268,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
||||||
})
|
})
|
||||||
.then((res) => {
|
.then((res) => {
|
||||||
if (controller.signal.aborted) return;
|
if (controller.signal.aborted) return;
|
||||||
set({
|
const isMultimodal = !!res.is_multimodal;
|
||||||
isDatasetMultimodal: !!res.is_multimodal,
|
const updates: Record<string, unknown> = {
|
||||||
|
isDatasetMultimodal: isMultimodal,
|
||||||
isCheckingDataset: false,
|
isCheckingDataset: false,
|
||||||
});
|
};
|
||||||
|
// train_on_responses_only: for VLMs, true if text dataset, false if vision dataset
|
||||||
|
if (get().isVisionModel) {
|
||||||
|
updates.trainOnCompletions = !isMultimodal;
|
||||||
|
}
|
||||||
|
set(updates as Partial<TrainingConfigState>);
|
||||||
})
|
})
|
||||||
.catch(() => {
|
.catch(() => {
|
||||||
if (controller.signal.aborted) return;
|
if (controller.signal.aborted) return;
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue