# Source: content/notes/deep-learning/transfer-learning-and-finetuning.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import copy
import torch
from torch import nn

torch.manual_seed(21)
torch.set_num_threads(1)
x = torch.randn(1024, 6)
source_y = (x[:, 0] + 0.5 * x[:, 1] > 0).long()
target_y = (x[:, 0] + 0.5 * x[:, 1] + 0.8 * x[:, 2] > 0).long()
base = nn.Sequential(nn.Linear(6, 24), nn.Tanh(), nn.Linear(24, 2))
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(base.parameters(), lr=0.02)
for _ in range(100):
    optimizer.zero_grad(set_to_none=True)
    loss = criterion(base(x[:640]), source_y[:640])
    loss.backward()
    optimizer.step()
base.eval()
base_snapshot = {k: v.detach().clone() for k, v in base.state_dict().items()}

class LowRankLinear(nn.Module):
    def __init__(self, linear, rank=2):
        super().__init__()
        self.base = linear
        self.base.requires_grad_(False)
        self.a = nn.Linear(linear.in_features, rank, bias=False)
        self.b = nn.Linear(rank, linear.out_features, bias=False)
        nn.init.zeros_(self.b.weight)
        self.scale = 2.0

    def forward(self, values):
        return self.base(values) + self.scale * self.b(self.a(values))

results = {}
for strategy in ("head", "full", "lora"):
    model = copy.deepcopy(base)
    if strategy != "full":
        model.requires_grad_(False)
        model[2].requires_grad_(True)
    if strategy == "lora":
        model[0] = LowRankLinear(model[0])
        torch.testing.assert_close(model(x[:8]), base(x[:8]))
        criterion(model(x[:96]), target_y[:96]).backward()
        assert model[0].a.weight.grad.eq(0).all()
        assert model[0].b.weight.grad.norm() > 0
        model.zero_grad(set_to_none=True)
    frozen_snapshot = {name: parameter.detach().clone()
                       for name, parameter in model.named_parameters()
                       if not parameter.requires_grad}
    params = [p for p in model.parameters() if p.requires_grad]
    opt = torch.optim.Adam(params, lr=0.01)
    initial = criterion(model(x[:96]), target_y[:96]).item()
    model.train()
    for _ in range(100):
        opt.zero_grad(set_to_none=True)
        loss = criterion(model(x[:96]), target_y[:96])
        loss.backward()
        opt.step()
    model.eval()
    with torch.no_grad():
        train_loss = criterion(model(x[:96]), target_y[:96]).item()
        predicted = model(x[768:]).argmax(1)
        results[strategy] = {
            "trainable": sum(p.numel() for p in params),
            "target_accuracy": (predicted == target_y[768:]).float().mean().item(),
            "source_retention": (predicted == source_y[768:]).float().mean().item(),
        }
        if strategy == "lora":
            layer = model[0]
            merged_weight = layer.base.weight + layer.scale * layer.b.weight @ layer.a.weight
            merged = nn.functional.linear(x[768:], merged_weight, layer.base.bias)
            torch.testing.assert_close(merged, layer(x[768:]), atol=2e-6, rtol=2e-5)
    assert train_loss < initial
    for name, parameter in model.named_parameters():
        if name in frozen_snapshot:
            torch.testing.assert_close(parameter, frozen_snapshot[name], rtol=0, atol=0)
for key, value in base.state_dict().items():
    torch.testing.assert_close(value, base_snapshot[key], rtol=0, atol=0)
assert results["lora"]["trainable"] < results["full"]["trainable"]
print(results)
