diff --git a/studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx b/studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx index 5117debcce..0a04a1e71f 100644 --- a/studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/hyperparameters-step.tsx @@ -23,7 +23,7 @@ import { TooltipTrigger, } from "@/components/ui/tooltip"; import { CONTEXT_LENGTHS } from "@/config/training"; -import { useTrainingConfigStore } from "@/features/training"; +import { useMaxStepsEpochsToggle, useTrainingConfigStore } from "@/features/training"; import { InformationCircleIcon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { useShallow } from "zustand/react/shallow"; @@ -35,6 +35,8 @@ export function HyperparametersStep() { setMaxSteps, epochs, setEpochs, + saveSteps, + setSaveSteps, contextLength, setContextLength, learningRate, @@ -52,6 +54,8 @@ export function HyperparametersStep() { setMaxSteps: s.setMaxSteps, epochs: s.epochs, setEpochs: s.setEpochs, + saveSteps: s.saveSteps, + setSaveSteps: s.setSaveSteps, contextLength: s.contextLength, setContextLength: s.setContextLength, learningRate: s.learningRate, @@ -67,6 +71,15 @@ export function HyperparametersStep() { const showLoraParams = trainingMethod === "lora" || trainingMethod === "qlora"; + const { useEpochs, toggleUseEpochs } = useMaxStepsEpochsToggle({ + maxSteps, + epochs, + saveSteps, + setMaxSteps, + setEpochs, + setSaveSteps, + }); + const maxStepsSliderMax = Math.max(500, maxSteps, 30); const epochsSliderMax = Math.max(10, epochs, 1); @@ -75,52 +88,84 @@ export function HyperparametersStep() {
Training
-
- - Max Steps - - - - - - Override total steps. Set 0 to use epochs instead.{" "} - - Read more - - - - -
- setMaxSteps(v)} - min={0} - max={maxStepsSliderMax} - step={1} - className="w-40" - /> - setMaxSteps(Number(e.target.value))} - min={0} - max={maxStepsSliderMax} - step={1} - className="w-16 text-right font-mono text-xs font-medium bg-muted/50 border border-border rounded-lg px-1.5 py-0.5 focus:outline-none focus:ring-1 focus:ring-primary/30 [&::-webkit-inner-spin-button]:appearance-none" - /> +
+
+ + {useEpochs ? "Epochs" : "Max Steps"} + + + + + + {useEpochs + ? "Number of full passes over the dataset." + : "Override total optimizer steps."}{" "} + + Read more + + + + +
+ + + useEpochs ? setEpochs(v) : setMaxSteps(v) + } + min={1} + max={useEpochs ? epochsSliderMax : maxStepsSliderMax} + step={1} + className="w-40" + /> + { + const raw = e.target.value; + if (raw === "") return; + + const value = Number(raw); + if (!Number.isFinite(value) || value < 1) return; + + if (useEpochs) { + setEpochs(value); + } else { + setMaxSteps(value); + } + }} + min={1} + max={useEpochs ? epochsSliderMax : maxStepsSliderMax} + step={1} + className="w-16 text-right font-mono text-xs font-medium bg-muted/50 border border-border rounded-lg px-1.5 py-0.5 focus:outline-none focus:ring-1 focus:ring-primary/30 [&::-webkit-inner-spin-button]:appearance-none" + /> +
@@ -206,55 +251,6 @@ export function HyperparametersStep() { />
-
- - Epochs - - - - - - Number of full passes over the dataset. Set 0 to run by max - steps.{" "} - - Read more - - - - -
- setEpochs(v)} - min={0} - max={epochsSliderMax} - step={1} - className="w-40" - /> - setEpochs(Number(e.target.value))} - min={0} - max={epochsSliderMax} - step={1} - className="w-12 text-right font-mono text-xs font-medium bg-muted/50 border border-border rounded-lg px-1.5 py-0.5 focus:outline-none focus:ring-1 focus:ring-primary/30 [&::-webkit-inner-spin-button]:appearance-none" - /> -
-
diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 5e7794a451..5758795206 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -29,7 +29,7 @@ import { OPTIMIZER_OPTIONS, TARGET_MODULES, } from "@/config/training"; -import { useTrainingConfigStore } from "@/features/training"; +import { useMaxStepsEpochsToggle, useTrainingConfigStore } from "@/features/training"; import type { GradientCheckpointing } from "@/types/training"; import { ArrowDown01Icon, @@ -120,6 +120,15 @@ export function ParamsSection(): ReactElement { const showVisionLora = store.isVisionModel && store.isDatasetImage === true; const [loraOpen, setLoraOpen] = useState(false); const [hyperOpen, setHyperOpen] = useState(false); + const { useEpochs, toggleUseEpochs } = useMaxStepsEpochsToggle({ + maxSteps: store.maxSteps, + epochs: store.epochs, + saveSteps: store.saveSteps, + setMaxSteps: store.setMaxSteps, + setEpochs: store.setEpochs, + setSaveSteps: store.setSaveSteps, + }); + const maxStepsSliderMax = Math.max(500, store.maxSteps, 30); const epochsSliderMax = Math.max(20, store.epochs, 1); @@ -133,56 +142,92 @@ export function ParamsSection(): ReactElement { className="md:min-h-[470px]" >
- {/* Max Steps */} + {/* Max Steps / Epochs */}
-
- - Max Steps - - - - - - Override total steps. Set 0 to use epochs instead.{" "} - - Read more - - - - - store.setMaxSteps(Number(e.target.value))} - min={0} - max={maxStepsSliderMax} +
+
+ + {useEpochs ? "Epochs" : "Max Steps"} + + + + + + {useEpochs + ? "Number of full passes over the dataset." + : "Override total optimizer steps."}{" "} + + Read more + + + + +
+ + { + const raw = e.target.value; + if (raw === "") return; + + const value = Number(raw); + if (!Number.isFinite(value) || value < 1) return; + + if (useEpochs) { + store.setEpochs(value); + } else { + store.setMaxSteps(value); + } + }} + min={1} + max={useEpochs ? epochsSliderMax : maxStepsSliderMax} + step={1} + className="w-16 text-right font-mono text-xs font-medium bg-muted/50 border border-border rounded-lg px-1.5 py-0.5 focus:outline-none focus:ring-1 focus:ring-primary/30 [&::-webkit-inner-spin-button]:appearance-none" + /> +
+
+ + useEpochs ? store.setEpochs(v) : store.setMaxSteps(v) + } + min={1} + max={useEpochs ? epochsSliderMax : maxStepsSliderMax} step={1} - className="w-16 text-right font-mono text-xs font-medium bg-muted/50 border border-border rounded-lg px-1.5 py-0.5 focus:outline-none focus:ring-1 focus:ring-primary/30 [&::-webkit-inner-spin-button]:appearance-none" /> +

+ {useEpochs + ? "Each epoch is one full pass over your dataset." + : "Limits training to a fixed number of optimizer steps."} +

- store.setMaxSteps(v)} - min={0} - max={maxStepsSliderMax} - step={1} - /> -

- Total optimizer steps. Use 0 to run by epochs. -

{/* Context length */} @@ -681,28 +726,30 @@ export function ParamsSection(): ReactElement { max={100} step={1} /> - - Number of full passes over the dataset. Set 0 to run by - max steps.{" "} - - Read more - - - } - value={store.epochs} - onChange={store.setEpochs} - min={0} - max={epochsSliderMax} - step={1} - /> + {!useEpochs && ( + + Number of full passes over the dataset. Set 0 to run by + max steps.{" "} + + Read more + + + } + value={store.epochs} + onChange={store.setEpochs} + min={0} + max={epochsSliderMax} + step={1} + /> + )} 0 ? value : DEFAULT_MAX_STEPS; +} + +function normalizePrevSaveSteps(value: number): number { + return Number.isFinite(value) && value >= 0 ? value : 0; +} + +type UseMaxStepsEpochsToggleParams = { + maxSteps: number; + epochs: number; + saveSteps: number; + setMaxSteps: (value: number) => void; + setEpochs: (value: number) => void; + setSaveSteps: (value: number) => void; + defaultEpochs?: number; +}; + +type UseMaxStepsEpochsToggleResult = { + useEpochs: boolean; + toggleUseEpochs: () => void; +}; + +export function useMaxStepsEpochsToggle({ + maxSteps, + epochs, + saveSteps, + setMaxSteps, + setEpochs, + setSaveSteps, + defaultEpochs = DEFAULT_EPOCHS, +}: UseMaxStepsEpochsToggleParams): UseMaxStepsEpochsToggleResult { + const useEpochs = maxSteps === 0; + const [prevMaxSteps, setPrevMaxSteps] = useState(() => + normalizePrevMaxSteps(readStoredNumber(PREV_MAX_STEPS_KEY, DEFAULT_MAX_STEPS)), + ); + const [prevSaveSteps, setPrevSaveSteps] = useState(() => { + if (maxSteps === 0 && saveSteps > 0) { + return normalizePrevSaveSteps(saveSteps); + } + return normalizePrevSaveSteps(readStoredNumber(PREV_SAVE_STEPS_KEY, 0)); + }); + + useEffect(() => { + if (maxSteps > 0) { + const normalized = normalizePrevMaxSteps(maxSteps); + setPrevMaxSteps(normalized); + writeStoredNumber(PREV_MAX_STEPS_KEY, normalized); + } + }, [maxSteps]); + + useEffect(() => { + if (!useEpochs) { + const normalized = normalizePrevSaveSteps(saveSteps); + setPrevSaveSteps(normalized); + writeStoredNumber(PREV_SAVE_STEPS_KEY, normalized); + } + }, [saveSteps, useEpochs]); + + const toggleUseEpochs = useCallback(() => { + if (useEpochs) { + setMaxSteps(normalizePrevMaxSteps(prevMaxSteps)); + setSaveSteps(normalizePrevSaveSteps(prevSaveSteps)); + return; + } + + setMaxSteps(0); + setEpochs(epochs || defaultEpochs); + setSaveSteps(0); + }, [ + defaultEpochs, + epochs, + prevMaxSteps, + prevSaveSteps, + setEpochs, + setMaxSteps, + setSaveSteps, + useEpochs, + ]); + + return { useEpochs, toggleUseEpochs }; +} diff --git a/studio/frontend/src/features/training/index.ts b/studio/frontend/src/features/training/index.ts index 7adddfd168..9a70e0a71c 100644 --- a/studio/frontend/src/features/training/index.ts +++ b/studio/frontend/src/features/training/index.ts @@ -8,6 +8,7 @@ export { } from "./stores/training-runtime-store"; export { useTrainingActions } from "./hooks/use-training-actions"; export { useTrainingRuntimeLifecycle } from "./hooks/use-training-runtime-lifecycle"; +export { useMaxStepsEpochsToggle } from "./hooks/use-max-steps-epochs-toggle"; export { HfDatasetSubsetSplitSelectors } from "./components/hf-dataset-subset-split-selectors"; export { useDatasetPreviewDialogStore } from "./stores/dataset-preview-dialog-store"; export { uploadTrainingDataset } from "./api/datasets-api";