import torch
import torch.nn as nn
from typing import Optional, List
class FrozenEmbedding(nn.Module):
"""Non-trainable embedding layer as training foundation."""
def __init__(self, vocab_size: int, embed_dim: int, initialize_random: bool = True):
super().__init__()
if initialize_random:
self.embed = nn.Embedding(vocab_size, embed_dim)
for param in self.embed.parameters():
param.requires_grad = False
else:
self.embed = nn.Embedding(vocab_size, embed_dim)
def forward(self, input_ids):
return self.embed(input_ids)
class TransformerBlock(nn.Module):
"""Single Transformer layer for independent training."""
def __init__(self, hidden_dim: int, num_heads: int, ff_dim: int, dropout: float = 0.1):
super().__init__()
self.attention = nn.MultiheadAttention(hidden_dim, num_heads, dropout=dropout, batch_first=True)
self.norm1 = nn.LayerNorm(hidden_dim)
self.norm2 = nn.LayerNorm(hidden_dim)
self.feed_forward = nn.Sequential(
nn.Linear(hidden_dim, ff_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(ff_dim, hidden_dim),
nn.Dropout(dropout)
)
def forward(self, x, attn_mask=None):
attn_out, _ = self.attention(x, x, x, attn_mask=attn_mask)
x = x + attn_out
x = self.norm1(x)
ff_out = self.feed_forward(x)
x = x + ff_out
x = self.norm2(x)
return x
class GrowingTransformer(nn.Module):
"""Model that grows layer-by-layer on frozen embeddings."""
def __init__(self, vocab_size: int, hidden_dim: int, num_heads: int, ff_dim: int):
super().__init__()
self.embedding = FrozenEmbedding(vocab_size, hidden_dim, initialize_random=True)
self.blocks = nn.ModuleList()
self.hidden_dim = hidden_dim
self.num_heads = num_heads
self.ff_dim = ff_dim
self.add_block()
def add_block(self):
"""Add a new trainable Transformer layer."""
new_block = TransformerBlock(self.hidden_dim, self.num_heads, self.ff_dim)
self.blocks.append(new_block)
def forward(self, input_ids):
x = self.embedding(input_ids)
for block in self.blocks:
x = block(x)
return x
def train_single_layer(model: GrowingTransformer, layer_idx: int, train_loader,
optimizer, criterion, num_epochs: int = 100,
patience: int = 5) -> float:
"""Train a single layer while freezing all others."""
for i, block in enumerate(model.blocks):
for param in block.parameters():
param.requires_grad = (i == layer_idx)
best_loss = float('inf')
patience_counter = 0
for epoch in range(num_epochs):
total_loss = 0
for batch_idx, (input_ids, target_ids) in enumerate(train_loader):
optimizer.zero_grad()
logits = model(input_ids)
logits = logits.view(-1, model.hidden_dim)
target = target_ids.view(-1)
loss = criterion(logits, target)
loss.backward()
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(train_loader)
if avg_loss < best_loss:
best_loss = avg_loss
patience_counter = 0
else:
patience_counter += 1
if patience_counter >= patience:
print(f"Layer {layer_idx} converged at epoch {epoch} with loss {best_loss:.4f}")
break
return best_loss
def sequential_growth(initial_model: GrowingTransformer, train_loader,
num_growth_steps: int = 5, epochs_per_layer: int = 100):
"""Grow model by repeatedly adding and training layers."""
for step in range(num_growth_steps):
print(f"\n=== Growth Step {step + 1}/{num_growth_steps} ===")
if step > 0:
initial_model.add_block()
optimizer = torch.optim.AdamW(initial_model.parameters(), lr=1e-4)
criterion = nn.CrossEntropyLoss()
layer_loss = train_single_layer(
initial_model,
layer_idx=step,
train_loader=train_loader,
optimizer=optimizer,
criterion=criterion,
num_epochs=epochs_per_layer,
patience=5
)
for param in initial_model.blocks[step].parameters():
param.requires_grad = False
print(f"Layer {step} frozen. Model now has {len(initial_model.blocks)} layers")
return initial_model
def apply_lora_finetuning(model: GrowingTransformer, train_loader, lr: float = 1e-3, epochs: int = 10):
"""Apply LoRA fine-tuning across all layers after growth."""
for block in model.blocks:
block.lora_q = nn.Linear(model.hidden_dim, 8)
block.lora_v = nn.Linear(8, model.hidden_dim)
optimizer = torch.optim.AdamW(
[p for p in model.parameters() if p.requires_grad],
lr=lr
)
for epoch in range(epochs):
total_loss = 0
for input_ids, target_ids in train_loader:
optimizer.zero_grad()
logits = model(input_ids)
loss = nn.CrossEntropyLoss()(logits.view(-1, model.hidden_dim), target_ids.view(-1))
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"LoRA Fine-tuning Epoch {epoch + 1}: Loss {total_loss / len(train_loader):.4f}")