S-Dreamer/CodeCraftLab
0
1import streamlit as st2import threading3import random4import time5from datetime import datetime6from utils import add_log, timestamp7 8# Handle missing dependencies9try:10 import torch11 import pandas as pd12 from transformers import TrainingArguments as HFTrainingArguments13 from transformers import Trainer, AutoModelForCausalLM, AutoTokenizer14 from datasets import Dataset, DatasetDict15 TRANSFORMERS_AVAILABLE = True16except ImportError:17 TRANSFORMERS_AVAILABLE = False18 HFTrainingArguments = None19 20 # For demo purposes21 class DummyTrainer:22 def __init__(self, **kwargs):23 self.callback = type('obj', (object,), {'__init__': lambda self: None})24 25 def train(self):26 pass27 28def initialize_training_progress(model_id):29 """30 Initialize training progress tracking for a model.31 32 Args:33 model_id: Identifier for the model34 """35 if 'training_progress' not in st.session_state:36 st.session_state.training_progress = {}37 38 st.session_state.training_progress[model_id] = {39 'status': 'initialized',40 'current_epoch': 0,41 'total_epochs': 0,42 'loss_history': [],43 'started_at': timestamp(),44 'completed_at': None,45 'progress': 0.046 }47 48def update_training_progress(model_id, epoch=None, loss=None, status=None, progress=None, total_epochs=None):49 """50 Update training progress for a model.51 52 Args:53 model_id: Identifier for the model54 epoch: Current epoch55 loss: Current loss value56 status: Training status57 progress: Progress percentage (0-100)58 total_epochs: Total number of epochs59 """60 if 'training_progress' not in st.session_state or model_id not in st.session_state.training_progress:61 initialize_training_progress(model_id)62 63 progress_data = st.session_state.training_progress[model_id]64 65 if epoch is not None:66 progress_data['current_epoch'] = epoch67 68 if loss is not None:69 progress_data['loss_history'].append(loss)70 71 if status is not None:72 progress_data['status'] = status73 if status == 'completed':74 progress_data['completed_at'] = timestamp()75 progress_data['progress'] = 100.076 77 if progress is not None:78 progress_data['progress'] = progress79 80 if total_epochs is not None:81 progress_data['total_epochs'] = total_epochs82 83def tokenize_dataset(dataset, tokenizer, max_length=512):84 """85 Tokenize a dataset for model training.86 87 Args:88 dataset: The dataset to tokenize89 tokenizer: The tokenizer to use90 max_length: Maximum sequence length91 92 Returns:93 Dataset: Tokenized dataset94 """95 def tokenize_function(examples):96 return tokenizer(examples['code'], padding='max_length', truncation=True, max_length=max_length)97 98 tokenized_dataset = dataset.map(tokenize_function, batched=True)99 return tokenized_dataset100 101def train_model_thread(model_id, dataset_name, base_model_name, training_args, device, stop_event):102 """103 Thread function for training a model.104 105 Args:106 model_id: Identifier for the model107 dataset_name: Name of the dataset to use108 base_model_name: Base model from Hugging Face109 training_args: Training arguments110 device: Device to use for training (cpu/cuda)111 stop_event: Threading event to signal stopping112 """113 try:114 # Get dataset115 dataset = st.session_state.datasets[dataset_name]['data']116 117 # Initialize model and tokenizer118 add_log(f"Initializing model {base_model_name} for training")119 tokenizer = AutoTokenizer.from_pretrained(base_model_name)120 model = AutoModelForCausalLM.from_pretrained(base_model_name)121 122 # Check if tokenizer has padding token123 if tokenizer.pad_token is None:124 tokenizer.pad_token = tokenizer.eos_token125 model.config.pad_token_id = model.config.eos_token_id126 127 # Tokenize dataset128 add_log(f"Tokenizing dataset {dataset_name}")129 train_dataset = tokenize_dataset(dataset['train'], tokenizer)130 val_dataset = tokenize_dataset(dataset['validation'], tokenizer)131 132 # Update training progress133 update_training_progress(134 model_id, 135 status='running',136 total_epochs=training_args.num_train_epochs137 )138 139 # Define custom callback to track progress140 class CustomCallback(Trainer.callback):141 def on_epoch_end(self, args, state, control, **kwargs):142 current_epoch = state.epoch143 epoch_loss = state.log_history[-1].get('loss', 0)144 update_training_progress(145 model_id, 146 epoch=current_epoch, 147 loss=epoch_loss,148 progress=(current_epoch / training_args.num_train_epochs) * 100149 )150 add_log(f"Epoch {current_epoch}/{training_args.num_train_epochs} completed. Loss: {epoch_loss:.4f}")151 152 # Check if training should be stopped153 if stop_event.is_set():154 add_log(f"Training for model {model_id} was manually stopped")155 control.should_training_stop = True156 157 # Configure training arguments158 args = HFTrainingArguments(159 output_dir=f"./results/{model_id}",160 evaluation_strategy="epoch",161 learning_rate=training_args.learning_rate,162 per_device_train_batch_size=training_args.batch_size,163 per_device_eval_batch_size=training_args.batch_size,164 num_train_epochs=training_args.num_train_epochs,165 weight_decay=0.01,166 save_total_limit=1,167 )168 169 # Initialize trainer170 trainer = Trainer(171 model=model,172 args=args,173 train_dataset=train_dataset,174 eval_dataset=val_dataset,175 tokenizer=tokenizer,176 callbacks=[CustomCallback]177 )178 179 # Train the model180 add_log(f"Starting training for model {model_id}")181 trainer.train()182 183 # Save the model184 if not stop_event.is_set():185 add_log(f"Training completed for model {model_id}")186 update_training_progress(model_id, status='completed')187 188 # Save to session state189 st.session_state.trained_models[model_id] = {190 'model': model,191 'tokenizer': tokenizer,192 'info': {193 'id': model_id,194 'base_model': base_model_name,195 'dataset': dataset_name,196 'created_at': timestamp(),197 'epochs': training_args.num_train_epochs,198 'learning_rate': training_args.learning_rate,199 'batch_size': training_args.batch_size200 }201 }202 203 except Exception as e:204 add_log(f"Error during training model {model_id}: {str(e)}", "ERROR")205 update_training_progress(model_id, status='failed')206 207class TrainingArguments:208 def __init__(self, learning_rate, batch_size, num_train_epochs):209 self.learning_rate = learning_rate210 self.batch_size = batch_size211 self.num_train_epochs = num_train_epochs212 213def start_model_training(model_id, dataset_name, base_model_name, learning_rate, batch_size, epochs):214 """215 Start model training in a separate thread.216 217 Args:218 model_id: Identifier for the model219 dataset_name: Name of the dataset to use220 base_model_name: Base model from Hugging Face221 learning_rate: Learning rate for training222 batch_size: Batch size for training223 epochs: Number of training epochs224 225 Returns:226 threading.Event: Event to signal stopping the training227 """228 # Use simulate_training instead if transformers isn't available229 if not TRANSFORMERS_AVAILABLE:230 add_log("No transformers library available, using simulation mode")231 return simulate_training(model_id, dataset_name, base_model_name, epochs)232 233 # Create training arguments234 training_args = TrainingArguments(235 learning_rate=learning_rate,236 batch_size=batch_size,237 num_train_epochs=epochs238 )239 240 # Determine device241 device = "cuda" if torch.cuda.is_available() else "cpu"242 add_log(f"Using device: {device}")243 244 # Initialize training progress245 initialize_training_progress(model_id)246 247 # Create stop event248 stop_event = threading.Event()249 250 # Start training thread251 training_thread = threading.Thread(252 target=train_model_thread,253 args=(model_id, dataset_name, base_model_name, training_args, device, stop_event)254 )255 training_thread.start()256 257 return stop_event258 259def stop_model_training(model_id, stop_event):260 """261 Stop model training.262 263 Args:264 model_id: Identifier for the model265 stop_event: Threading event to signal stopping266 """267 if stop_event.is_set():268 return269 270 add_log(f"Stopping training for model {model_id}")271 stop_event.set()272 273 # Update training progress274 if 'training_progress' in st.session_state and model_id in st.session_state.training_progress:275 progress_data = st.session_state.training_progress[model_id]276 if progress_data['status'] == 'running':277 progress_data['status'] = 'stopped'278 progress_data['completed_at'] = timestamp()279 280def get_running_training_jobs():281 """282 Get list of currently running training jobs.283 284 Returns:285 list: List of model IDs with running training jobs286 """287 running_jobs = []288 289 if 'training_progress' in st.session_state:290 for model_id, progress in st.session_state.training_progress.items():291 if progress['status'] == 'running':292 running_jobs.append(model_id)293 294 return running_jobs295 296# For demo purposes - Simulate training progress without actual model training297def simulate_training_thread(model_id, dataset_name, base_model_name, epochs, stop_event):298 """299 Simulate training progress for demonstration purposes.300 301 Args:302 model_id: Identifier for the model303 dataset_name: Name of the dataset to use304 base_model_name: Base model from Hugging Face305 epochs: Number of training epochs306 stop_event: Threading event to signal stopping307 """308 add_log(f"Starting simulated training for model {model_id}")309 update_training_progress(model_id, status='running', total_epochs=epochs)310 311 for epoch in range(1, epochs + 1):312 if stop_event.is_set():313 add_log(f"Simulated training for model {model_id} was manually stopped")314 update_training_progress(model_id, status='stopped')315 return316 317 # Simulate epoch time318 time.sleep(2)319 320 # Generate random loss that decreases over time321 loss = max(0.1, 2.0 - (epoch / epochs) * 1.5 + random.uniform(-0.1, 0.1))322 323 # Update progress324 update_training_progress(325 model_id,326 epoch=epoch,327 loss=loss,328 progress=(epoch / epochs) * 100329 )330 331 add_log(f"Epoch {epoch}/{epochs} completed. Loss: {loss:.4f}")332 333 # Training completed334 add_log(f"Simulated training completed for model {model_id}")335 update_training_progress(model_id, status='completed')336 337 # Create dummy model and tokenizer338 tokenizer = AutoTokenizer.from_pretrained(base_model_name)339 model = AutoModelForCausalLM.from_pretrained(base_model_name)340 341 # Save to session state342 st.session_state.trained_models[model_id] = {343 'model': model,344 'tokenizer': tokenizer,345 'info': {346 'id': model_id,347 'base_model': base_model_name,348 'dataset': dataset_name,349 'created_at': timestamp(),350 'epochs': epochs,351 'simulated': True352 }353 }354 355def simulate_training(model_id, dataset_name, base_model_name, epochs):356 """357 Start simulated training in a separate thread.358 359 Args:360 model_id: Identifier for the model361 dataset_name: Name of the dataset to use362 base_model_name: Base model from Hugging Face363 epochs: Number of training epochs364 365 Returns:366 threading.Event: Event to signal stopping the training367 """368 # Initialize training progress369 initialize_training_progress(model_id)370 371 # Create stop event372 stop_event = threading.Event()373 374 # Start training thread375 training_thread = threading.Thread(376 target=simulate_training_thread,377 args=(model_id, dataset_name, base_model_name, epochs, stop_event)378 )379 training_thread.start()380 381 return stop_event382 