CoolFace
Apppublic

S-Dreamer/CodeCraftLab

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
training_utils.py382 linesDownload Raw Back to root
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