{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torch-geometric\n!pip install numpy\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T09:50:29.746587Z","iopub.execute_input":"2025-12-19T09:50:29.747252Z","iopub.status.idle":"2025-12-19T09:50:39.471398Z","shell.execute_reply.started":"2025-12-19T09:50:29.747221Z","shell.execute_reply":"2025-12-19T09:50:39.470528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== TRANSFORMER ON 50K SAMPLES - FIXED VERSION ====================\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix\nfrom sklearn.preprocessing import StandardScaler\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport time\nimport os\nimport warnings\nwarnings.filterwarnings('ignore')\nfrom tqdm import tqdm\nfrom torch.nn.utils.rnn import pad_sequence\n\nprint(\"=\"*80)\nprint(\"⚡ TRANSFORMER TRAINING ON 50,000 SAMPLES - FIXED VERSION\")\nprint(\"=\"*80)\n\n# 1. DATA LOADING\nprint(\"\\n📂 LOADING DATA...\")\ncsv_path = '/kaggle/input/hms-harmful-brain-activity-classification/train.csv'\ntrain_csv = pd.read_csv(csv_path)\nprint(f\"✅ Loaded {len(train_csv):,} samples\")\n\n# Process labels\nlabel_cols = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nvotes = train_csv[label_cols].values.astype(float)\nrow_sums = votes.sum(axis=1, keepdims=True)\nrow_sums[row_sums == 0] = 1\nprobabilities = votes / row_sums\ntrain_csv['dominant_class'] = np.argmax(probabilities, axis=1)\n\n# Select 50K samples\nn_samples = 50000\nif len(train_csv) > n_samples:\n    _, sample_indices = train_test_split(\n        range(len(train_csv)),\n        test_size=n_samples,\n        random_state=42,\n        stratify=train_csv['dominant_class']\n    )\n    selected_df = train_csv.iloc[sample_indices]\nelse:\n    selected_df = train_csv\n    n_samples = len(selected_df)\n\nprint(f\"Selected {n_samples:,} samples\")\n\n# 2. SEQUENCE DATA EXTRACTION - FIXED\nprint(\"\\n🔧 EXTRACTING SEQUENCE DATA...\")\n\nDATA_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs'\n\ndef extract_sequence_data(eeg_data):\n    \"\"\"Extract sequence data for Transformer - FIXED VERSION\"\"\"\n    if eeg_data.size == 0:\n        return np.zeros((100, 4), dtype=np.float32) + 1e-8\n    \n    # Downsample to 100 timesteps\n    if eeg_data.shape[0] > 100:\n        indices = np.linspace(0, eeg_data.shape[0]-1, 100).astype(int)\n        eeg_data = eeg_data[indices, :]\n    \n    # Use first 4 channels\n    if eeg_data.shape[1] > 4:\n        eeg_data = eeg_data[:, :4]\n    elif eeg_data.shape[1] < 4:\n        eeg_data = np.pad(eeg_data, ((0, 0), (0, 4 - eeg_data.shape[1])), mode='constant', constant_values=1e-8)\n    \n    # SAFER Normalization per channel with robust statistics\n    eps = 1e-6\n    for ch in range(eeg_data.shape[1]):\n        channel_data = eeg_data[:, ch]\n        \n        # Handle constant channels\n        if np.std(channel_data) < eps:\n            eeg_data[:, ch] = np.random.normal(0, eps, size=channel_data.shape)\n        else:\n            # Robust scaling using median and IQR\n            median_val = np.median(channel_data)\n            q75, q25 = np.percentile(channel_data, [75, 25])\n            iqr = q75 - q25 + eps\n            \n            if iqr < eps:\n                # If IQR is too small, use standard deviation\n                scale = np.std(channel_data) + eps\n                eeg_data[:, ch] = (channel_data - median_val) / scale\n            else:\n                eeg_data[:, ch] = (channel_data - median_val) / iqr\n    \n    # Clip extreme values to prevent outliers\n    eeg_data = np.clip(eeg_data, -5, 5)\n    \n    # Final check for NaN/Inf\n    eeg_data = np.nan_to_num(eeg_data, nan=0.0, posinf=5.0, neginf=-5.0)\n    \n    return eeg_data.astype(np.float32)\n\n# Process samples\nsequence_list = []\nlabels_list = []\nfailed = 0\n\nprocessing_start = time.time()\n\nfor idx in tqdm(range(n_samples), desc=\"Processing EEGs\"):\n    try:\n        sample_id = selected_df.iloc[idx]['eeg_id']\n        eeg_path = f\"{DATA_DIR}/{sample_id}.parquet\"\n        \n        if os.path.exists(eeg_path):\n            eeg_df = pd.read_parquet(eeg_path)\n            if not eeg_df.empty and eeg_df.shape[0] > 10:\n                eeg_data = eeg_df.values.astype(np.float32)\n                sequence = extract_sequence_data(eeg_data)\n                sequence_list.append(sequence)\n                labels_list.append(selected_df.iloc[idx]['dominant_class'])\n            else:\n                # Create dummy data for missing samples\n                sequence = np.random.normal(0, 0.1, (100, 4)).astype(np.float32)\n                sequence_list.append(sequence)\n                labels_list.append(selected_df.iloc[idx]['dominant_class'])\n                failed += 1\n        else:\n            # Create dummy data for missing files\n            sequence = np.random.normal(0, 0.1, (100, 4)).astype(np.float32)\n            sequence_list.append(sequence)\n            labels_list.append(selected_df.iloc[idx]['dominant_class'])\n            failed += 1\n    except Exception as e:\n        # Create dummy data on error\n        sequence = np.random.normal(0, 0.1, (100, 4)).astype(np.float32)\n        sequence_list.append(sequence)\n        labels_list.append(selected_df.iloc[idx]['dominant_class'])\n        failed += 1\n\nprocessing_time = time.time() - processing_start\nprint(f\"\\n✅ Processed {len(sequence_list):,}/{n_samples} sequences\")\nprint(f\"   Failed: {failed:,}\")\nprint(f\"   Time: {processing_time:.1f}s\")\n\n# Check data statistics\nprint(\"\\n📊 DATA STATISTICS:\")\nsample_seq = sequence_list[0]\nprint(f\"  Sequence shape: {sample_seq.shape}\")\nprint(f\"  Min: {np.min([s.min() for s in sequence_list]):.4f}\")\nprint(f\"  Max: {np.max([s.max() for s in sequence_list]):.4f}\")\nprint(f\"  Mean: {np.mean([s.mean() for s in sequence_list]):.4f}\")\nprint(f\"  Std: {np.mean([s.std() for s in sequence_list]):.4f}\")\nprint(f\"  Has NaN: {any(np.any(np.isnan(s)) for s in sequence_list)}\")\nprint(f\"  Has Inf: {any(np.any(np.isinf(s)) for s in sequence_list)}\")\n\n# 3. DATA PREPARATION\ny = np.array(labels_list, dtype=np.int32)\n\n# Check label distribution\nunique_labels, label_counts = np.unique(y, return_counts=True)\nprint(\"\\n🎯 LABEL DISTRIBUTION:\")\nfor label, count in zip(unique_labels, label_counts):\n    print(f\"  Class {label}: {count:,} samples ({count/len(y)*100:.1f}%)\")\n\n# Split data\ntrain_idx, val_idx = train_test_split(\n    range(len(sequence_list)),\n    test_size=0.1,\n    random_state=42,\n    stratify=y\n)\n\nprint(f\"\\n📊 DATASET SIZES:\")\nprint(f\"   Training: {len(train_idx):,} sequences\")\nprint(f\"   Validation: {len(val_idx):,} sequences\")\n\n# 4. TRANSFORMER MODEL - FIXED VERSION\nclass EEG_Transformer(nn.Module):\n    def __init__(self, input_channels=4, d_model=64, nhead=4, num_layers=3, num_classes=6):\n        super().__init__()\n        \n        # Input projection with layer normalization\n        self.input_proj = nn.Linear(input_channels, d_model)\n        self.ln1 = nn.LayerNorm(d_model)\n        \n        # Positional encoding\n        self.pos_encoder = PositionalEncoding(d_model)\n        \n        # Transformer layers\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model, \n            nhead=nhead, \n            dim_feedforward=256,  # Increased\n            dropout=0.1,  # Reduced dropout\n            activation='relu',  # Changed to relu for stability\n            batch_first=True,\n            norm_first=True  # Important for stability\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        \n        # Multi-scale pooling\n        self.avg_pool = nn.AdaptiveAvgPool1d(1)\n        self.max_pool = nn.AdaptiveMaxPool1d(1)\n        \n        # Classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(d_model * 2, 128),\n            nn.LayerNorm(128),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            nn.Linear(128, 64),\n            nn.LayerNorm(64),\n            nn.ReLU(),\n            nn.Dropout(0.1),\n            nn.Linear(64, num_classes)\n        )\n        \n        # Initialize weights properly\n        self._init_weights()\n    \n    def _init_weights(self):\n        \"\"\"Initialize weights properly\"\"\"\n        for p in self.parameters():\n            if p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n    \n    def forward(self, x):\n        # Input projection\n        x = self.input_proj(x)\n        x = self.ln1(x)\n        \n        # Add positional encoding\n        x = self.pos_encoder(x)\n        \n        # Transformer\n        x = self.transformer(x)\n        \n        # Multi-scale pooling\n        x = x.transpose(1, 2)  # [batch, features, seq_len]\n        avg_pooled = self.avg_pool(x).squeeze(-1)\n        max_pooled = self.max_pool(x).squeeze(-1)\n        pooled = torch.cat([avg_pooled, max_pooled], dim=1)\n        \n        # Classifier\n        logits = self.classifier(pooled)\n        \n        return logits  # Return raw logits\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=5000):\n        super().__init__()\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe.unsqueeze(0))\n    \n    def forward(self, x):\n        return x + self.pe[:, :x.size(1)]\n\n# 5. TRAINING SETUP\nprint(f\"\\n{'='*60}\")\nprint(\"⚡ TRAINING TRANSFORMER\")\nprint(f\"{'='*60}\")\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Create datasets\nclass SequenceDataset(torch.utils.data.Dataset):\n    def __init__(self, sequences, labels):\n        self.sequences = sequences\n        self.labels = labels\n    \n    def __len__(self):\n        return len(self.sequences)\n    \n    def __getitem__(self, idx):\n        return torch.FloatTensor(self.sequences[idx]), self.labels[idx]\n\ndef collate_fn(batch):\n    sequences = [item[0] for item in batch]\n    labels = [item[1] for item in batch]\n    sequences_padded = pad_sequence(sequences, batch_first=True)\n    return sequences_padded, torch.tensor(labels, dtype=torch.long)\n\ntrain_dataset = SequenceDataset(\n    [sequence_list[i] for i in train_idx],\n    [labels_list[i] for i in train_idx]\n)\nval_dataset = SequenceDataset(\n    [sequence_list[i] for i in val_idx],\n    [labels_list[i] for i in val_idx]\n)\n\ntrain_loader = torch.utils.data.DataLoader(\n    train_dataset, batch_size=32, shuffle=True, collate_fn=collate_fn, num_workers=2\n)\nval_loader = torch.utils.data.DataLoader(\n    val_dataset, batch_size=32, shuffle=False, collate_fn=collate_fn, num_workers=2\n)\n\n# Initialize model\nmodel = EEG_Transformer(input_channels=4, d_model=64, nhead=4, num_layers=2, num_classes=6)  # Reduced layers\nmodel = model.to(device)\n\n# Use CrossEntropyLoss (which includes softmax)\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.1)  # Added label smoothing\n\n# AdamW optimizer with weight decay\noptimizer = torch.optim.AdamW(\n    model.parameters(), \n    lr=0.001,  # Increased learning rate\n    weight_decay=0.01,\n    betas=(0.9, 0.999),\n    eps=1e-8\n)\n\n# Cosine annealing scheduler\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, \n    T_max=15,  # Matches epochs\n    eta_min=1e-6\n)\n\n# Track metrics\ntrain_losses = []\nval_losses = []\ntrain_accuracies = []\nval_accuracies = []\nlearning_rates = []\n\nnum_epochs = 15\nbest_val_acc = 0\npatience_counter = 0\nmax_patience = 5\n\nprint(f\"\\nTraining for {num_epochs} epochs...\")\nprint(f\"Batch size: 32\")\nprint(f\"Learning rate: {optimizer.param_groups[0]['lr']}\")\nprint(f\"Loss function: CrossEntropyLoss with label smoothing\")\n\nstart_time = time.time()\n\n# Training loop\nfor epoch in range(num_epochs):\n    # Training phase\n    model.train()\n    train_loss = 0\n    train_correct = 0\n    train_total = 0\n    \n    train_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Train]\", leave=False)\n    \n    for batch_idx, (sequences, labels) in enumerate(train_bar):\n        sequences, labels = sequences.to(device), labels.to(device)\n        \n        # Forward pass with gradient clipping\n        optimizer.zero_grad()\n        outputs = model(sequences)\n        loss = criterion(outputs, labels)\n        \n        # Check for NaN in loss\n        if torch.isnan(loss):\n            print(f\"⚠️  NaN detected in training loss at batch {batch_idx}\")\n            loss = torch.tensor(1.0, requires_grad=True).to(device)\n        \n        loss.backward()\n        \n        # Gradient clipping\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        \n        optimizer.step()\n        \n        # Statistics\n        train_loss += loss.item()\n        _, predicted = torch.max(outputs.data, 1)\n        train_total += labels.size(0)\n        train_correct += (predicted == labels).sum().item()\n        \n        # Update progress bar\n        train_bar.set_postfix({\n            'loss': loss.item(),\n            'acc': train_correct / train_total\n        })\n    \n    avg_train_loss = train_loss / len(train_loader)\n    train_acc = train_correct / train_total\n    train_losses.append(avg_train_loss)\n    train_accuracies.append(train_acc)\n    \n    # Validation phase\n    model.eval()\n    val_loss = 0\n    val_correct = 0\n    val_total = 0\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        val_bar = tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Val]\", leave=False)\n        \n        for sequences, labels in val_bar:\n            sequences, labels = sequences.to(device), labels.to(device)\n            outputs = model(sequences)\n            loss = criterion(outputs, labels)\n            \n            val_loss += loss.item()\n            _, predicted = torch.max(outputs.data, 1)\n            val_total += labels.size(0)\n            val_correct += (predicted == labels).sum().item()\n            \n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            \n            val_bar.set_postfix({\n                'loss': loss.item(),\n                'acc': val_correct / val_total\n            })\n    \n    avg_val_loss = val_loss / len(val_loader)\n    val_acc = val_correct / val_total\n    val_losses.append(avg_val_loss)\n    val_accuracies.append(val_acc)\n    \n    # Update scheduler\n    scheduler.step()\n    learning_rates.append(optimizer.param_groups[0]['lr'])\n    \n    print(f\"\\n  Epoch {epoch+1} Summary:\")\n    print(f\"    Train Loss: {avg_train_loss:.4f}, Train Acc: {train_acc:.4f}\")\n    print(f\"    Val Loss: {avg_val_loss:.4f}, Val Acc: {val_acc:.4f}\")\n    print(f\"    LR: {learning_rates[-1]:.6f}\")\n    \n    # Save best model\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        patience_counter = 0\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'val_acc': val_acc,\n            'train_acc': train_acc,\n        }, 'best_transformer_model.pth')\n        print(f\"    💾 Saved best model with val_acc: {val_acc:.4f}\")\n    else:\n        patience_counter += 1\n        if patience_counter >= max_patience:\n            print(f\"    ⏹️  Early stopping at epoch {epoch+1}\")\n            break\n\ntotal_time = time.time() - start_time\n\n# Load best model\ncheckpoint = torch.load('best_transformer_model.pth')\nmodel.load_state_dict(checkpoint['model_state_dict'])\n\n# Final evaluation\nmodel.eval()\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for sequences, labels in val_loader:\n        sequences, labels = sequences.to(device), labels.to(device)\n        outputs = model(sequences)\n        _, predicted = torch.max(outputs.data, 1)\n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\nfinal_val_acc = accuracy_score(all_labels, all_preds)\n\nprint(f\"\\n✅ Transformer Training Complete!\")\nprint(f\"   Best Validation Accuracy: {best_val_acc:.4f}\")\nprint(f\"   Final Validation Accuracy: {final_val_acc:.4f}\")\nprint(f\"   Total Training Time: {total_time:.1f}s\")\nprint(f\"   Epochs completed: {len(train_losses)}\")\n\n# Save final model\ntorch.save(model.state_dict(), 'transformer_50k_model_final.pth')\nprint(\"   Model saved as 'transformer_50k_model_final.pth'\")\n\n# 6. CLASSIFICATION REPORT\nprint(f\"\\n{'='*60}\")\nprint(\"📊 CLASSIFICATION REPORT\")\nprint(f\"{'='*60}\")\n\nclass_names = ['Seizure', 'LPD', 'GPD', 'LRDA', 'GRDA', 'Other']\n\nprint(\"\\n📋 Detailed Classification Report:\")\nreport = classification_report(all_labels, all_preds, target_names=class_names, digits=4)\nprint(report)\n\n# Confusion Matrix\ncm = confusion_matrix(all_labels, all_preds)\n\nplt.figure(figsize=(12, 10))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n            xticklabels=class_names, yticklabels=class_names)\nplt.title('Transformer - Confusion Matrix (50K Samples)', fontsize=14, fontweight='bold')\nplt.xlabel('Predicted Label', fontsize=12)\nplt.ylabel('True Label', fontsize=12)\nplt.tight_layout()\nplt.savefig('transformer_confusion_matrix.png', dpi=300, bbox_inches='tight')\nplt.show()\n\n# 7. TRAINING CURVES\nprint(f\"\\n{'='*60}\")\nprint(\"📈 TRAINING CURVES\")\nprint(f\"{'='*60}\")\n\nfig, axes = plt.subplots(2, 2, figsize=(15, 12))\n\n# Loss Curves\nepochs_range = range(1, len(train_losses) + 1)\n\nax1 = axes[0, 0]\nax1.plot(epochs_range, train_losses, 'o-', linewidth=2, markersize=6, label='Training', color='blue')\nax1.plot(epochs_range, val_losses, 's-', linewidth=2, markersize=6, label='Validation', color='red')\nax1.set_xlabel('Epoch', fontsize=12)\nax1.set_ylabel('Loss', fontsize=12)\nax1.set_title('Loss Curves', fontsize=14, fontweight='bold')\nax1.grid(True, alpha=0.3)\nax1.legend()\n\n# Accuracy Curves\nax2 = axes[0, 1]\nax2.plot(epochs_range, train_accuracies, 'o-', linewidth=2, markersize=6, label='Training', color='blue')\nax2.plot(epochs_range, val_accuracies, 's-', linewidth=2, markersize=6, label='Validation', color='red')\nax2.set_xlabel('Epoch', fontsize=12)\nax2.set_ylabel('Accuracy', fontsize=12)\nax2.set_title('Accuracy Curves', fontsize=14, fontweight='bold')\nax2.grid(True, alpha=0.3)\nax2.legend()\n\n# Learning Rate\nax3 = axes[1, 0]\nax3.plot(epochs_range, learning_rates, 'o-', linewidth=2, markersize=6, color='green')\nax3.set_xlabel('Epoch', fontsize=12)\nax3.set_ylabel('Learning Rate', fontsize=12)\nax3.set_title('Learning Rate Schedule', fontsize=14, fontweight='bold')\nax3.grid(True, alpha=0.3)\nax3.set_yscale('log')\n\n# Class Distribution\nax4 = axes[1, 1]\nclass_counts = [np.sum(all_labels == i) for i in range(6)]\nax4.bar(class_names, class_counts, color='skyblue', edgecolor='black')\nax4.set_xlabel('Class', fontsize=12)\nax4.set_ylabel('Count', fontsize=12)\nax4.set_title('Validation Set Class Distribution', fontsize=14, fontweight='bold')\nax4.grid(True, alpha=0.3, axis='y')\n\nfor i, count in enumerate(class_counts):\n    ax4.text(i, count + max(class_counts)*0.01, str(count), ha='center', fontsize=10)\n\nplt.suptitle('Transformer Training Analysis - 50,000 Samples', fontsize=16, fontweight='bold')\nplt.tight_layout()\nplt.savefig('transformer_training_curves.png', dpi=300, bbox_inches='tight')\nplt.show()\n\n# 8. PERFORMANCE SUMMARY\nprint(f\"\\n{'='*60}\")\nprint(\"📊 PERFORMANCE SUMMARY\")\nprint(f\"{'='*60}\")\n\nprint(f\"\\n🎯 Final Results:\")\nprint(f\"   Best Validation Accuracy: {best_val_acc:.4f}\")\nprint(f\"   Final Validation Accuracy: {final_val_acc:.4f}\")\nprint(f\"   Best Epoch: {checkpoint['epoch'] + 1}\")\nprint(f\"   Training Accuracy at Best Epoch: {checkpoint['train_acc']:.4f}\")\n\nprint(f\"\\n⏱️  Timing:\")\nprint(f\"   Data Processing: {processing_time:.1f}s\")\nprint(f\"   Model Training: {total_time:.1f}s\")\nprint(f\"   Total Time: {processing_time + total_time:.1f}s\")\nprint(f\"   Time per Epoch: {total_time/len(train_losses):.1f}s\")\n\nprint(f\"\\n📈 Model Insights:\")\nprint(f\"   1. Transformer with {sum(p.numel() for p in model.parameters()):,} parameters\")\nprint(f\"   2. Input shape: {sample_seq.shape}\")\nprint(f\"   3. Using {device} for training\")\n\n# Calculate class-wise accuracy\nprint(f\"\\n🎯 Class-wise Performance:\")\nfor i, class_name in enumerate(class_names):\n    class_mask = np.array(all_labels) == i\n    if np.sum(class_mask) > 0:\n        class_acc = np.mean(np.array(all_preds)[class_mask] == np.array(all_labels)[class_mask])\n        print(f\"   {class_name}: {class_acc:.4f} ({np.sum(class_mask)} samples)\")\n\nprint(f\"\\n✅ Transformer training completed successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T09:50:39.473618Z","iopub.execute_input":"2025-12-19T09:50:39.473970Z","iopub.status.idle":"2025-12-19T10:52:20.090300Z","shell.execute_reply.started":"2025-12-19T09:50:39.473926Z","shell.execute_reply":"2025-12-19T10:52:20.089544Z"}},"outputs":[],"execution_count":null}]}