help@rskworld.in +91 93305 39277
RSK World
  • Home
  • Development
    • Web Development
    • Mobile Apps
    • Software
    • Games
    • Project
  • Technologies
    • Data Science
    • AI Development
    • Cloud Development
    • Blockchain
    • Cyber Security
    • Dev Tools
    • Testing Tools
  • Blog
  • About
  • Contact

Theme Settings

Color Scheme
Display Options
Font Size
100%
Back to Project
RSK World
pytorch-neuralnetworks
/
examples
RSK World
pytorch-neuralnetworks
Neural networks with PyTorch
examples
  • advanced_features_example.py5.1 KB
  • hyperparameter_tuning_example.py4.6 KB
  • transfer_learning_example.py2.1 KB
advanced_features_example.py
examples/advanced_features_example.py
Raw Download
Find: Go to:
"""
Advanced Features Example - PyTorch Neural Networks
Project: PyTorch Neural Networks
Author: RSK World
Website: https://rskworld.in
Email: help@rskworld.in
Phone: +91 93305 39277
Description: Example demonstrating advanced training features
"""

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset
import sys
import os

# Add parent directory to path
sys.path.append('..')

from models.basic_nn import BasicNeuralNetwork
from training.advanced_trainer import AdvancedTrainer
from training.callbacks import EarlyStopping, ModelCheckpoint, LearningRateScheduler
from training.metrics import evaluate_model, plot_confusion_matrix, classification_report_metrics
from training.utils import generate_sample_data
from utils.tensorboard_logger import TensorBoardLogger


def advanced_training_example():
    """
    Demonstrate advanced training features
    """
    print("=" * 60)
    print("Advanced Training Features Example")
    print("Author: RSK World (https://rskworld.in)")
    print("=" * 60)
    
    # Set device
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"\nUsing device: {device}\n")
    
    # Generate data
    print("Generating sample data...")
    X_train, y_train, X_val, y_val = generate_sample_data(
        n_samples=1000, n_features=20, n_classes=3
    )
    X_test, y_test, _, _ = generate_sample_data(
        n_samples=200, n_features=20, n_classes=3
    )
    
    # Create data loaders
    train_loader = DataLoader(TensorDataset(X_train, y_train), batch_size=32, shuffle=True)
    val_loader = DataLoader(TensorDataset(X_val, y_val), batch_size=32, shuffle=False)
    test_loader = DataLoader(TensorDataset(X_test, y_test), batch_size=32, shuffle=False)
    
    # Create model
    model = BasicNeuralNetwork(input_size=20, hidden_size=64, output_size=3).to(device)
    
    # Define loss and optimizer
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    
    # Create learning rate scheduler
    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)
    lr_scheduler = LearningRateScheduler(scheduler)
    
    # Create callbacks
    early_stopping = EarlyStopping(patience=5, verbose=True)
    checkpoint = ModelCheckpoint(
        filepath='../saved_models/best_model.pth',
        monitor='val_loss',
        save_best_only=True
    )
    
    # Create TensorBoard logger
    logger = TensorBoardLogger(log_dir='../runs/advanced_example')
    
    # Create advanced trainer with gradient clipping
    trainer = AdvancedTrainer(
        model, criterion, optimizer, device,
        gradient_clip=1.0,  # Clip gradients to norm 1.0
        use_mixed_precision=False  # Set to True if using CUDA
    )
    
    # Training loop with callbacks
    num_epochs = 20
    best_val_loss = float('inf')
    
    print("\nStarting training with advanced features...")
    print("-" * 60)
    
    for epoch in range(num_epochs):
        # Train
        train_loss, train_acc = trainer.train_epoch(train_loader)
        
        # Validate
        val_results = evaluate_model(model, val_loader, device, criterion)
        val_loss = val_results['loss']
        val_acc = val_results['accuracy']
        
        # Update learning rate
        lr_scheduler.step(val_loss)
        
        # Log metrics
        logger.log_training_metrics({'loss': train_loss, 'accuracy': train_acc}, epoch)
        logger.log_validation_metrics({'loss': val_loss, 'accuracy': val_acc}, epoch)
        logger.log_learning_rate(optimizer, epoch)
        logger.increment_step()
        
        print(f"Epoch {epoch+1}/{num_epochs}")
        print(f"Train - Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%")
        print(f"Val   - Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%")
        
        # Checkpoint
        checkpoint(val_loss, model, optimizer, epoch)
        
        # Early stopping
        if early_stopping(val_loss, model):
            print("Early stopping triggered!")
            break
        
        print("-" * 60)
    
    logger.close()
    
    # Final evaluation on test set
    print("\nEvaluating on test set...")
    test_results = evaluate_model(model, test_loader, device, criterion)
    print(f"Test Accuracy: {test_results['accuracy']:.2f}%")
    
    # Confusion matrix
    plot_confusion_matrix(
        test_results['labels'],
        test_results['predictions'],
        class_names=['Class 0', 'Class 1', 'Class 2'],
        save_path='../confusion_matrix.png'
    )
    
    # Classification report
    print("\nClassification Report:")
    classification_report_metrics(
        test_results['labels'],
        test_results['predictions'],
        class_names=['Class 0', 'Class 1', 'Class 2']
    )
    
    print("\n" + "=" * 60)
    print("Example completed!")
    print("Check TensorBoard: tensorboard --logdir=runs")
    print("=" * 60)


if __name__ == '__main__':
    advanced_training_example()

156 lines•5.1 KB
python
🚀 Support RSK World

Subscribe to our YouTube channel for latest tutorials & updates!



Click subscribe & support our work ❤️

About RSK World

Founded by Molla Samser, with Designer & Tester Rima Khatun, RSK World is your one-stop destination for free programming resources, source code, and development tools.

Founder: Molla Samser
Designer & Tester: Rima Khatun

Development

  • Game Development
  • Web Development
  • Mobile Development
  • AI Development
  • Development Tools

Legal

  • Terms & Conditions
  • Privacy Policy
  • Disclaimer

Contact Info

Nutanhat, Mongolkote
Purba Burdwan, West Bengal
India, 713147

+91 93305 39277

hello@rskworld.in
support@rskworld.in

© 2026 RSK World. All rights reserved.

Content used for educational purposes only. View Disclaimer