Model Collapse: When Artificial Intelligence Gets Sick From Its Own Data
Imagine photocopying a photo, then photocopying that copy, and repeating the process ten times in a row. What happens? Each time, fine details fadeβ¦
Training Error << Validation ErrorExample:Training Accuracy = 99%Validation Accuracy = 65% β π₯ Overfitting!
Simple line passing through points- Too simple- High training error- High validation error
Smooth curve capturing general pattern- Appropriate complexity- Low training error- Low validation error
Complex curve passing exactly through all points- Too complex- Training error = 0- Validation error very high!
# Real data: y = 2x + 1 + noisex_train = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10]y_train = [3.1, 5.2, 6.8, 9.1, 10.9, 13.2, 15.1, 16.8, 19.2, 20.9]# Model 1: Degree 1 (simple line) - Underfitting# y = ax + b# Training Error = 0.5# Model 2: Degree 2 (curve) - Good Fit# y = axΒ² + bx + c# Training Error = 0.1# Validation Error = 0.12 β Good!# Model 3: Degree 9 (too complex) - Overfitting# y = aβxβΉ + aβxβΈ + ... + aβx + aβ# Training Error = 0.0001 β Looks great# Validation Error = 15.7 β π₯ Disaster!
# Example: Cat vs Dog classification with 10 training images# Simple model: 100 parameters βmodel_simple = Sequential([Dense(10, activation='relu'),Dense(2, activation='softmax')])# Complex model: 10,000,000 parameters βmodel_complex = Sequential([Dense(1000, activation='relu'),Dense(1000, activation='relu'),Dense(1000, activation='relu'),Dense(1000, activation='relu'),Dense(2, activation='softmax')])
Number of parameters << Number of training samplesGood: 1000 parameters, 100,000 samplesBad: 1,000,000 parameters, 100 samples
# Scenario 1: 1000 images# ResNet-50 (25M parameters)# Result: Severe overfitting! β# Scenario 2: 1,000,000 images# ResNet-50 (25M parameters)# Result: Works β# Scenario 3: 1000 images + Data Augmentation# ResNet-50 (25M parameters)# Result: Better β
# Epoch 1: Train=80%, Val=78% β Good# Epoch 10: Train=95%, Val=90% β Great# Epoch 20: Train=98%, Val=92% β Best point!# Epoch 50: Train=99.5%, Val=88% β Starting to overfit# Epoch 100: Train=99.9%, Val=75% β π₯ Complete overfitting!
# Training data:# 90% correct labels# 10% incorrect labels (noise)# Very powerful model:# Learns to fit even the noise!# Result: Performs poorly on real data (without noise)
# Example: Face recognition# All training images: bright light, frontal angle# Model: Only learns these conditions# Test: Image with low light, different angle# Result: Failure! β
Total Error = BiasΒ² + Variance + Irreducible ErrorIrreducible Error = Inherent noise in data (uncontrollable)
| Feature | High Bias (Underfitting) | Sweet Spot | High Variance (Overfitting) |
|---|---|---|---|
| Training Error | High (e.g., 30%) | Low (e.g., 5%) | Very low (e.g., 0.1%) |
| Validation Error | High (e.g., 32%) | Low (e.g., 6%) | Very high (e.g., 25%) |
| Gap | Small (2%) | Small (1%) | Very large (24.9%) |
| Model Complexity | Too simple | Appropriate | Too complex |
| Solution | More complex model, more features | - | Regularization, more data |
def diagnose_model(train_acc, val_acc):gap = train_acc - val_accif train_acc < 0.7 and val_acc < 0.7:return "High Bias (Underfitting) - Model too simple!"elif gap > 0.15: # Gap more than 15%return "High Variance (Overfitting) - Model is memorizing!"elif train_acc > 0.9 and val_acc > 0.85:return "Sweet Spot - Excellent! π"else:return "Needs more investigation"# Example:print(diagnose_model(0.99, 0.65)) # High Variance (Overfitting)print(diagnose_model(0.65, 0.63)) # High Bias (Underfitting)print(diagnose_model(0.92, 0.89)) # Sweet Spot
Loss_total = Loss_data + Ξ» Ξ£(wΒ²)Ξ» = regularization coefficient (usually 0.001 to 0.1)
# Method 1: In optimizeroptimizer = torch.optim.Adam(model.parameters(),lr=0.001,weight_decay=0.01 # This is L2!)# Method 2: Manual in lossdef l2_regularization(model, lambda_=0.01):l2_loss = 0for param in model.parameters():l2_loss += torch.sum(param ** 2)return lambda_ * l2_loss# In training loop:loss = criterion(output, target) + l2_regularization(model)
Loss_total = Loss_data + Ξ» Ξ£|w|def l1_regularization(model, lambda_=0.01):l1_loss = 0for param in model.parameters():l1_loss += torch.sum(torch.abs(param))return lambda_ * l1_lossloss = criterion(output, target) + l1_regularization(model)
class ModelWithDropout(nn.Module):def __init__(self):super().__init__()self.fc1 = nn.Linear(784, 512)self.dropout1 = nn.Dropout(p=0.5) # 50% neurons turned offself.fc2 = nn.Linear(512, 256)self.dropout2 = nn.Dropout(p=0.3) # 30% offself.fc3 = nn.Linear(256, 10)def forward(self, x):x = F.relu(self.fc1(x))x = self.dropout1(x) # Dropout only in trainingx = F.relu(self.fc2(x))x = self.dropout2(x)x = self.fc3(x)return x# Important: Turn off dropout in evaluationmodel.eval() # Automatically turns off dropout
class EarlyStopping:def __init__(self, patience=7, min_delta=0.001):"""patience: How many epochs to waitmin_delta: Minimum acceptable improvement"""self.patience = patienceself.min_delta = min_deltaself.counter = 0self.best_loss = Noneself.early_stop = Falseself.best_model = Nonedef __call__(self, val_loss, model):if self.best_loss is None:self.best_loss = val_lossself.best_model = copy.deepcopy(model.state_dict())elif val_loss > self.best_loss - self.min_delta:# Validation loss didn't improveself.counter += 1print(f'EarlyStopping counter: {self.counter}/{self.patience}')if self.counter >= self.patience:self.early_stop = Trueelse:# Validation loss improvedself.best_loss = val_lossself.best_model = copy.deepcopy(model.state_dict())self.counter = 0return self.early_stopdef load_best_model(self, model):"""Load best model"""model.load_state_dict(self.best_model)# Usage:early_stopping = EarlyStopping(patience=10, min_delta=0.001)for epoch in range(100):train_loss = train_epoch()val_loss = validate()print(f'Epoch {epoch}: Train Loss={train_loss:.4f}, Val Loss={val_loss:.4f}')if early_stopping(val_loss, model):print("Early stopping triggered!")break# Load best modelearly_stopping.load_best_model(model)
from torchvision import transformstrain_transform = transforms.Compose([transforms.RandomHorizontalFlip(p=0.5), # Horizontal fliptransforms.RandomRotation(degrees=15), # Rotate Β±15 degreestransforms.ColorJitter( # Color changebrightness=0.2,contrast=0.2,saturation=0.2),transforms.RandomCrop(224, padding=4), # Random croptransforms.ToTensor(),transforms.Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225])])train_dataset = ImageFolder(train_dir, transform=train_transform)
def augment_text(text):"""Augmentation techniques for text"""augmented = []# 1. Synonym Replacementaugmented.append(replace_with_synonyms(text))# 2. Random Deletionaugmented.append(random_deletion(text, p=0.1))# 3. Random Swapaugmented.append(random_swap(text))# 4. Back Translationaugmented.append(back_translate(text, target_lang='fr'))return augmented
from sklearn.ensemble import BaggingClassifierfrom sklearn.tree import DecisionTreeClassifier# Train 10 different models on different data subsetsbagging = BaggingClassifier(base_estimator=DecisionTreeClassifier(),n_estimators=10,max_samples=0.8, # Each model sees 80% of databootstrap=True)bagging.fit(X_train, y_train)predictions = bagging.predict(X_test)
class EnsembleModel:def __init__(self, models):self.models = modelsdef predict(self, x):predictions = []for model in self.models:model.eval()with torch.no_grad():pred = model(x)predictions.append(pred)# Average predictionsensemble_pred = torch.mean(torch.stack(predictions), dim=0)return ensemble_pred# Train 5 models with different initializationmodels = []for i in range(5):model = create_model()train(model) # Train with different seedmodels.append(model)# Use ensembleensemble = EnsembleModel(models)final_prediction = ensemble.predict(test_data)
# Combination of techniques:class ResNetBlock(nn.Module):def __init__(self):# 1. Residual Connections for gradient flowself.conv1 = nn.Conv2d(...)# 2. Batch Normalizationself.bn1 = nn.BatchNorm2d(...)# 3. Dropout (in later versions)self.dropout = nn.Dropout(0.2)def forward(self, x):residual = xout = self.conv1(x)out = self.bn1(out)out = F.relu(out)out = self.dropout(out)out += residual # Skip connectionreturn out# 4. Heavy Data Augmentationtrain_transform = transforms.Compose([transforms.RandomResizedCrop(224),transforms.RandomHorizontalFlip(),transforms.ColorJitter(0.4, 0.4, 0.4, 0.1),transforms.ToTensor()])# 5. Weight Decayoptimizer = SGD(model.parameters(), lr=0.1, weight_decay=1e-4)
# 1. Huge data (45TB text!)# 2. Dropout everywhereclass GPT3Layer:def __init__(self):self.attention_dropout = nn.Dropout(0.1)self.residual_dropout = nn.Dropout(0.1)self.output_dropout = nn.Dropout(0.1)# 3. Weight Decayoptimizer = AdamW(model.parameters(), weight_decay=0.1)# 4. Gradient Clippingtorch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)# 5. Learning Rate Scheduling with Warmup# 6. Early Stopping based on validation loss