| import os |
| import torch |
| import torch.nn.functional as F |
| from transformers import AutoTokenizer |
| from pathlib import Path |
| import logging |
| from tqdm import tqdm |
| import json |
| from datetime import datetime |
| from model import MultiModalDenseTransformer |
| from data_loader import create_pretrain_dataloader |
|
|
| logging.basicConfig( |
| level=logging.INFO, |
| format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' |
| ) |
| logger = logging.getLogger(__name__) |
| os.environ["HF_ENDPOINT"] = "https://hf-mirror.com" |
|
|
|
|
| class PreTrainer: |
| def __init__( |
| self, |
| model: MultiModalDenseTransformer, |
| tokenizer, |
| learning_rate: float = 3e-4, |
| weight_decay: float = 0.1, |
| warmup_steps: int = 1000, |
| max_steps: int = 100000, |
| gradient_accumulation_steps: int = 16, |
| max_grad_norm: float = 1.0, |
| log_interval: int = 10, |
| save_interval: int = 1000, |
| checkpoint_dir: str = "checkpoints/pretrain", |
| loss_log_file: str = "checkpoints/pretrain/train_loss.log" |
| ): |
| self.model = model |
| self.tokenizer = tokenizer |
| self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| |
| self.model.to(self.device) |
| |
| self.optimizer = torch.optim.AdamW( |
| model.parameters(), |
| lr=learning_rate, |
| weight_decay=weight_decay, |
| betas=(0.9, 0.95), |
| eps=1e-8 |
| ) |
| |
| from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts |
| |
| self.warmup_steps = warmup_steps |
| self.max_lr = learning_rate |
| self.min_lr = learning_rate * 0.1 |
| self.current_step = 0 |
| |
| self.use_amp = torch.cuda.is_available() |
| self.scaler = torch.amp.GradScaler('cuda', enabled=self.use_amp) |
| |
| self.gradient_accumulation_steps = gradient_accumulation_steps |
| self.max_grad_norm = max_grad_norm |
| self.max_steps = max_steps |
| self.log_interval = log_interval |
| self.save_interval = save_interval |
| |
| self.checkpoint_dir = Path(checkpoint_dir) |
| self.checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| |
| self.loss_log_file = Path(loss_log_file) |
| self.loss_log_file.parent.mkdir(parents=True, exist_ok=True) |
| |
| self.global_step = 0 |
| self.tokens_seen = 0 |
| self.running_loss = 0.0 |
| self.best_loss = float('inf') |
| |
| logger.info(f"PreTrainer initialized:") |
| logger.info(f" Device: {self.device}") |
| logger.info(f" Learning Rate: {learning_rate}") |
| logger.info(f" Max Steps: {max_steps}") |
| logger.info(f" Gradient Accumulation: {gradient_accumulation_steps}") |
| logger.info(f" Effective Batch Size: {gradient_accumulation_steps}") |
| logger.info(f" Mixed Precision: {self.use_amp}") |
|
|
| def _get_lr(self) -> float: |
| if self.current_step < self.warmup_steps: |
| return self.max_lr * (self.current_step / self.warmup_steps) |
| else: |
| progress = (self.current_step - self.warmup_steps) / (self.max_steps - self.warmup_steps) |
| return self.min_lr + (self.max_lr - self.min_lr) * 0.5 * (1 + torch.cos(torch.tensor(progress * 3.14159))) |
|
|
| def _set_lr(self, lr: float): |
| for param_group in self.optimizer.param_groups: |
| param_group['lr'] = lr |
|
|
| def train_step(self, batch: dict) -> dict: |
| input_ids = batch['input_ids'].to(self.device) |
| attention_mask = batch['attention_mask'].to(self.device) |
| batch_size, seq_len = input_ids.shape |
| position_ids= torch.zeros_like(input_ids) |
| |
| for i in range(batch_size): |
| non_pad_mask = attention_mask[i].bool() |
| if non_pad_mask.any(): |
| positions = torch.cumsum(non_pad_mask.long(), dim=0) -1 |
| position_ids[i]=positions * non_pad_mask.long() |
|
|
| input_data = { |
| 'segments': [{ |
| 'type': 'text', |
| 'data': input_ids, |
| 'modality_id': 0 |
| }] |
| } |
| |
| with torch.amp.autocast('cuda', enabled=self.use_amp): |
| outputs = self.model( |
| input_data, |
| attention_mask=attention_mask, |
| position_ids=position_ids) |
| logits = outputs['logits'] |
| |
| shift_logits = logits[:, :-1, :].contiguous() |
| shift_labels = input_ids[:, 1:].contiguous() |
| shift_attention_mask = attention_mask[:, 1:].contiguous() |
| |
| loss = F.cross_entropy( |
| shift_logits.view(-1, shift_logits.size(-1)), |
| shift_labels.view(-1), |
| reduction='none' |
| ) |
| |
| loss = (loss * shift_attention_mask.view(-1)).sum() / (shift_attention_mask.sum() + 1e-8) |
| loss_for_backward = loss / self.gradient_accumulation_steps |
| |
| self.scaler.scale(loss_for_backward).backward() |
| self.tokens_seen += attention_mask.sum().item() |
| |
| return { |
| 'loss': loss.item(), |
| 'lr': self.optimizer.param_groups[0]['lr'] |
| } |
|
|
| def optimizer_step(self): |
| self.scaler.unscale_(self.optimizer) |
| |
| grad_norm = torch.nn.utils.clip_grad_norm_( |
| self.model.parameters(), |
| self.max_grad_norm |
| ) |
| |
| self.scaler.step(self.optimizer) |
| self.scaler.update() |
| self.optimizer.zero_grad(set_to_none=True) |
| |
| self.current_step += 1 |
| self.global_step += 1 |
| lr = self._get_lr() |
| self._set_lr(lr) |
| |
| return grad_norm.item() |
|
|
| def _write_loss_to_txt(self, step, avg_loss, lr, tokens_seen): |
| log_content = ( |
| f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] " |
| f"Step: {step}/{self.max_steps}, " |
| f"Average Loss: {avg_loss:.4f}, " |
| f"Learning Rate: {lr:.2e}, " |
| f"Tokens Seen: {tokens_seen/1e9:.2f}B\n" |
| ) |
| with open(self.loss_log_file, 'a', encoding='utf-8') as f: |
| f.write(log_content) |
|
|
| def train(self, dataloader, resume_from=None): |
| if resume_from: |
| self.load_checkpoint(resume_from) |
| if not self.loss_log_file.exists(): |
| with open(self.loss_log_file, 'w', encoding='utf-8') as f: |
| f.write("Training Log (Real Loss Values)\n") |
| f.write("="*80 + "\n") |
| |
| self.model.train() |
| progress_bar = tqdm(total=self.max_steps, initial=self.global_step) |
| |
| step_in_accumulation = 0 |
| accumulated_loss = 0.0 |
| |
| batches_to_skip = self.global_step * self.gradient_accumulation_steps |
| |
| logger.info(f"Current Global Step: {self.global_step}") |
| if batches_to_skip > 0: |
| logger.info(f"Resuming: Need to skip {batches_to_skip} batches to restore data state...") |
| |
| |
| data_iterator = iter(dataloader) |
| skipped = 0 |
| if batches_to_skip > 0: |
| with tqdm(total=batches_to_skip, desc="Skipping trained batches", unit="batch") as skip_pbar: |
| while skipped < batches_to_skip: |
| try: |
| _ = next(data_iterator) |
| skipped += 1 |
| skip_pbar.update(1) |
| except StopIteration: |
| logger.error("Dataset exhausted during skipping!") |
| return |
|
|
| try: |
| while True: |
| try: |
| batch = next(data_iterator) |
| except StopIteration: |
| break |
|
|
| if batch is None or batch['input_ids'].size(0) == 0: |
| continue |
| stats = self.train_step(batch) |
| step_in_accumulation += 1 |
| accumulated_loss += stats['loss'] |
| |
| if step_in_accumulation >= self.gradient_accumulation_steps: |
| avg_step_loss = accumulated_loss / self.gradient_accumulation_steps |
| |
| grad_norm = self.optimizer_step() |
| stats['grad_norm'] = grad_norm |
| stats['loss'] = avg_step_loss |
| |
| self.running_loss += avg_step_loss |
| |
| step_in_accumulation = 0 |
| accumulated_loss = 0.0 |
| |
| progress_bar.update(1) |
| progress_bar.set_postfix({ |
| 'loss': f"{stats['loss']:.4f}", |
| 'lr': f"{stats['lr']:.2e}", |
| 'tokens': f"{self.tokens_seen/1e9:.2f}B", |
| 'grad': f"{grad_norm:.2f}" |
| }) |
| |
| if self.global_step % self.log_interval == 0: |
| avg_loss = self.running_loss / self.log_interval |
| |
| logger.info( |
| f"Step {self.global_step}/{self.max_steps} | " |
| f"Loss: {avg_loss:.4f} | " |
| f"LR: {stats['lr']:.2e} | " |
| f"GradNorm: {grad_norm:.2f} | " |
| f"Tokens: {self.tokens_seen/1e9:.2f}B" |
| ) |
| |
| if avg_loss < self.best_loss: |
| self.best_loss = avg_loss |
| logger.info(f" New best loss: {self.best_loss:.4f}") |
| |
| self._write_loss_to_txt( |
| step=self.global_step, |
| avg_loss=avg_loss, |
| lr=stats['lr'], |
| tokens_seen=self.tokens_seen |
| ) |
| self.running_loss = 0.0 |
| if self.global_step % self.save_interval == 0: |
| self.save_checkpoint( |
| self.checkpoint_dir / f"step_{self.global_step}.pt" |
| ) |
| |
| if self.global_step >= self.max_steps: |
| break |
| |
| except KeyboardInterrupt: |
| logger.info("\n Training interrupted by user") |
| self.save_checkpoint( |
| self.checkpoint_dir / f"interrupted_step_{self.global_step}.pt" |
| ) |
| |
| finally: |
| progress_bar.close() |
|
|
| self.save_checkpoint(self.checkpoint_dir / "final_model.pt") |
|
|
| def save_checkpoint(self, path: Path): |
| checkpoint = { |
| 'model_state_dict': self.model.state_dict(), |
| 'optimizer_state_dict': self.optimizer.state_dict(), |
| 'scaler_state_dict': self.scaler.state_dict() if self.use_amp else None, |
| 'global_step': self.global_step, |
| 'current_step': self.current_step, |
| 'tokens_seen': self.tokens_seen, |
| 'best_loss': self.best_loss, |
| 'timestamp': datetime.now().isoformat() |
| } |
| |
| torch.save(checkpoint, path) |
| logger.info(f" Checkpoint saved to {path}") |
|
|
| def load_checkpoint(self, path: str): |
| checkpoint = torch.load(path, map_location=self.device, weights_only=True) |
| |
| self.model.load_state_dict(checkpoint['model_state_dict']) |
| self.optimizer.load_state_dict(checkpoint['optimizer_state_dict']) |
| |
| if self.use_amp and checkpoint.get('scaler_state_dict'): |
| self.scaler.load_state_dict(checkpoint['scaler_state_dict']) |
| |
| self.global_step = checkpoint['global_step'] |
| self.current_step = checkpoint.get('current_step', self.global_step) |
| self.tokens_seen = checkpoint['tokens_seen'] |
| self.best_loss = checkpoint.get('best_loss', float('inf')) |
|
|
|
|
| def main(): |
| config = { |
| 'model_dim': 1536, |
| 'vocab_size': 151665, |
| 'n_layers': 12, |
| 'n_heads': 12, |
| 'n_kv_heads': 4, |
| 'max_seq_len': 1024, |
| 'dropout': 0.1, |
| 'use_moe': False, |
| |
| |
| 'batch_size': 4, |
| 'gradient_accumulation_steps': 8, |
| 'learning_rate': 1e-4, |
| 'weight_decay': 0.1, |
| 'warmup_steps': 500, |
| 'max_steps': 100000, |
| 'max_grad_norm': 1.0, |
| |
| 'data_mix': 'skypile_training', |
| 'max_length': 1024, |
| 'num_workers': 2, |
| |
| 'log_interval': 10, |
| 'save_interval': 5000, |
| 'checkpoint_dir': 'checkpoints/pretrain_fixed', |
| 'loss_log_file': 'checkpoints/pretrain_fixed/train_loss_skypile_training.log' |
| } |
|
|
| tokenizer = AutoTokenizer.from_pretrained( |
| "Qwen/Qwen2.5-7B-Instruct", |
| use_fast=True, |
| trust_remote_code=True |
| ) |
| |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
| tokenizer.pad_token_id = tokenizer.eos_token_id |
| |
| config['vocab_size'] = len(tokenizer) |
|
|
| logger.info("Initializing model...") |
| model = MultiModalDenseTransformer( |
| model_dim=config['model_dim'], |
| vocab_size=config['vocab_size'], |
| n_layers=config['n_layers'], |
| n_heads=config['n_heads'], |
| n_kv_heads=config['n_kv_heads'], |
| max_seq_len=config['max_seq_len'], |
| dropout=config['dropout'], |
| use_moe=config['use_moe'], |
| use_gradient_checkpointing=True, |
| rope_scaling_type="yarn", |
| use_multimodal_fusion=False, |
| use_contrastive=False |
| ) |
|
|
| dataloader = create_pretrain_dataloader( |
| mix_name=config['data_mix'], |
| tokenizer=tokenizer, |
| batch_size=config['batch_size'], |
| num_workers=config['num_workers'], |
| max_length=config['max_length'] |
| ) |
| trainer = PreTrainer( |
| model=model, |
| tokenizer=tokenizer, |
| learning_rate=config['learning_rate'], |
| weight_decay=config['weight_decay'], |
| warmup_steps=config['warmup_steps'], |
| max_steps=config['max_steps'], |
| gradient_accumulation_steps=config['gradient_accumulation_steps'], |
| max_grad_norm=config['max_grad_norm'], |
| log_interval=config['log_interval'], |
| save_interval=config['save_interval'], |
| checkpoint_dir=config['checkpoint_dir'], |
| loss_log_file=config['loss_log_file'] |
| ) |
|
|
| logger.info("\n Starting fresh training with fixes...\n") |
| trainer.train(dataloader, resume_from="/root/checkpoints/pretrain_fixed/step_35000.pt") |
| |
|
|
|
|
| if __name__ == "__main__": |
| main() |