From 0b1c635b433c5a0169364203c42fd67aadc86f34 Mon Sep 17 00:00:00 2001 From: sshah229 Date: Sun, 15 Feb 2026 05:35:32 -0700 Subject: [PATCH 1/2] resolved the prgress metrics --- studio/backend/core/training/trainer.py | 60 ++++++++++++++++++++++++- studio/backend/routes/training.py | 35 ++++++++++----- 2 files changed, 83 insertions(+), 12 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index 5a3d8683e2..c77460406b 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -16,6 +16,7 @@ import json import threading import math import logging +import time from typing import Optional, Callable from dataclasses import dataclass import pandas as pd @@ -46,6 +47,10 @@ class TrainingProgress: is_completed: bool = False error: Optional[str] = None status_message: str = "Ready to train" # Current stage message + elapsed_seconds: Optional[float] = None + eta_seconds: Optional[float] = None + grad_norm: Optional[float] = None + num_tokens: Optional[int] = None class UnslothTrainer: """ @@ -67,6 +72,12 @@ class UnslothTrainer: self.is_vlm = False self.model_name = None + # Training metrics tracking + self.training_start_time: Optional[float] = None + self.batch_size: Optional[int] = None + self.max_seq_length: Optional[int] = None + self.gradient_accumulation_steps: Optional[int] = None + # Thread safety self._lock = threading.Lock() @@ -462,6 +473,14 @@ class UnslothTrainer: def _train_worker(self, dataset: Dataset, **training_args): """Worker function for training (runs in separate thread)""" try: + # Store training parameters for metrics calculation + self.batch_size = training_args.get('batch_size', 2) + self.max_seq_length = training_args.get('max_seq_length', 2048) + self.gradient_accumulation_steps = training_args.get('gradient_accumulation_steps', 4) + + # Set training start time + self.training_start_time = time.time() + self._update_progress(is_training=True, error=None) # Setup logging @@ -700,11 +719,50 @@ class UnslothTrainer: if logs: # Get loss from either 'loss' or 'train_loss' key loss_value = logs.get('loss', logs.get('train_loss', 0.0)) + current_step = state.global_step + + # Extract grad_norm from logs (available when gradient clipping is enabled) + grad_norm = logs.get('grad_norm', None) + + # Calculate elapsed_seconds + elapsed_seconds = None + if self.trainer_instance.training_start_time is not None: + elapsed_seconds = time.time() - self.trainer_instance.training_start_time + + # Calculate eta_seconds + eta_seconds = None + if elapsed_seconds is not None and current_step > 0: + total_steps = self.trainer_instance.training_progress.total_steps + if total_steps > 0: + steps_remaining = total_steps - current_step + if steps_remaining > 0: + time_per_step = elapsed_seconds / current_step + eta_seconds = time_per_step * steps_remaining + + # Calculate num_tokens + num_tokens = None + if (self.trainer_instance.batch_size is not None and + self.trainer_instance.max_seq_length is not None and + self.trainer_instance.gradient_accumulation_steps is not None): + # Approximate: batch_size * gradient_accumulation_steps * max_seq_length * steps + # Note: This is an approximation as actual sequences may be shorter + # Gradient accumulation means each step processes batch_size * gradient_accumulation_steps samples + num_tokens = ( + self.trainer_instance.batch_size * + self.trainer_instance.gradient_accumulation_steps * + self.trainer_instance.max_seq_length * + current_step + ) + self.trainer_instance._update_progress( - step=state.global_step, + step=current_step, epoch=round(state.epoch, 2) if state.epoch else 0, # Round epoch to 2 decimals loss=loss_value, learning_rate=logs.get('learning_rate', 0.0), + elapsed_seconds=elapsed_seconds, + eta_seconds=eta_seconds, + grad_norm=grad_norm, + num_tokens=num_tokens, status_message="" # Clear status message so metrics show ) diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 1a904cd371..e924d5329c 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -5,7 +5,7 @@ import sys from pathlib import Path from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse -from typing import Dict, Optional +from typing import Dict, Optional, Any import logging import asyncio from datetime import datetime @@ -480,6 +480,7 @@ async def stream_training_progress( learning_rate: float, total_steps: int, epoch: Optional[float] = None, + progress: Optional[Any] = None, ) -> TrainingProgress: total = max(total_steps, 0) if step < 0 or total == 0: @@ -489,6 +490,12 @@ async def stream_training_progress( float(step) / float(total) * 100.0 if total > 0 else 0.0 ) + # Get actual values from progress object if available + elapsed_seconds = getattr(progress, 'elapsed_seconds', None) if progress else None + eta_seconds = getattr(progress, 'eta_seconds', None) if progress else None + grad_norm = getattr(progress, 'grad_norm', None) if progress else None + num_tokens = getattr(progress, 'num_tokens', None) if progress else None + return TrainingProgress( job_id=job_id, step=step, @@ -497,10 +504,10 @@ async def stream_training_progress( learning_rate=learning_rate, progress_percent=progress_percent, epoch=epoch, - elapsed_seconds=None, - eta_seconds=None, - grad_norm=None, - num_tokens=None, + elapsed_seconds=elapsed_seconds, + eta_seconds=eta_seconds, + grad_norm=grad_norm, + num_tokens=num_tokens, ) def format_sse( @@ -534,7 +541,7 @@ async def stream_training_progress( ) total_replay = getattr(tp_replay, "total_steps", step_val) if tp_replay else step_val epoch_replay = getattr(tp_replay, "epoch", None) if tp_replay else None - payload = build_progress(step_val, loss_val, lr_val, total_replay, epoch_replay) + payload = build_progress(step_val, loss_val, lr_val, total_replay, epoch_replay, progress=tp_replay) yield format_sse(payload.model_dump_json(), event="progress", event_id=step_val) replayed += 1 if replayed: @@ -553,6 +560,7 @@ async def stream_training_progress( learning_rate=0.0, total_steps=initial_total_steps, epoch=initial_epoch, + progress=tp, ) yield format_sse(initial_progress.model_dump_json(), event="progress", event_id=0) @@ -566,11 +574,11 @@ async def stream_training_progress( getattr(tp, "total_steps", final_step) if tp else final_step ) final_epoch = getattr(tp, "epoch", None) if tp else None - payload = build_progress(final_step, final_loss, final_lr, final_total_steps, final_epoch) + payload = build_progress(final_step, final_loss, final_lr, final_total_steps, final_epoch, progress=tp) yield format_sse(payload.model_dump_json(), event="complete", event_id=final_step) else: yield format_sse( - build_progress(-1, 0.0, 0.0, 0).model_dump_json(), + build_progress(-1, 0.0, 0.0, 0, progress=tp).model_dump_json(), event="complete", event_id=0, ) @@ -605,6 +613,7 @@ async def stream_training_progress( current_lr, current_total_steps, current_epoch, + progress=tp_inner, ) yield format_sse( progress_payload.model_dump_json(), @@ -623,6 +632,7 @@ async def stream_training_progress( current_lr, current_total_steps, current_epoch, + progress=tp_inner, ) yield format_sse( heartbeat_payload.model_dump_json(), @@ -644,7 +654,7 @@ async def stream_training_progress( if tp_prep else 0 ) preparing_payload = build_progress( - 0, 0.0, 0.0, prep_total, + 0, 0.0, 0.0, prep_total, progress=tp_prep, ) yield format_sse( preparing_payload.model_dump_json(), @@ -655,7 +665,8 @@ async def stream_training_progress( # Timeout check if no_update_count > max_no_updates: logger.warning("Progress stream timeout - no updates received") - timeout_payload = build_progress(last_step, 0.0, 0.0, 0) + tp_timeout = getattr(getattr(backend, "trainer", None), "training_progress", None) + timeout_payload = build_progress(last_step, 0.0, 0.0, 0, progress=tp_timeout) yield format_sse( timeout_payload.model_dump_json(), event="error", @@ -667,7 +678,8 @@ async def stream_training_progress( except Exception as e: logger.error(f"Error in progress stream: {e}", exc_info=True) - error_payload = build_progress(0, 0.0, 0.0, 0) + tp_error = getattr(getattr(backend, "trainer", None), "training_progress", None) + error_payload = build_progress(0, 0.0, 0.0, 0, progress=tp_error) yield format_sse( error_payload.model_dump_json(), event="error", @@ -692,6 +704,7 @@ async def stream_training_progress( final_lr, final_total_steps, final_epoch, + progress=final_tp, ) yield format_sse( final_payload.model_dump_json(), From 63b34660edd7881252a913e1ffc84566fea8b0c3 Mon Sep 17 00:00:00 2001 From: sshah229 Date: Mon, 16 Feb 2026 00:38:40 -0700 Subject: [PATCH 2/2] modified the num_tokens logic --- studio/backend/core/training/trainer.py | 18 ++++-------------- 1 file changed, 4 insertions(+), 14 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index c77460406b..64f89af2a0 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -567,6 +567,7 @@ class UnslothTrainer: "seed": training_args.get('random_seed', 3407), "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 } # Add warmup parameter - use warmup_ratio if provided, otherwise warmup_steps @@ -739,20 +740,9 @@ class UnslothTrainer: time_per_step = elapsed_seconds / current_step eta_seconds = time_per_step * steps_remaining - # Calculate num_tokens - num_tokens = None - if (self.trainer_instance.batch_size is not None and - self.trainer_instance.max_seq_length is not None and - self.trainer_instance.gradient_accumulation_steps is not None): - # Approximate: batch_size * gradient_accumulation_steps * max_seq_length * steps - # Note: This is an approximation as actual sequences may be shorter - # Gradient accumulation means each step processes batch_size * gradient_accumulation_steps samples - num_tokens = ( - self.trainer_instance.batch_size * - self.trainer_instance.gradient_accumulation_steps * - self.trainer_instance.max_seq_length * - current_step - ) + # Extract num_tokens from TRL SFTTrainer state (real counter) + # Requires include_num_input_tokens_seen=True in SFTConfig + num_tokens = getattr(state, "num_input_tokens_seen", None) self.trainer_instance._update_progress( step=current_step,