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

import torch
from torch import nn

torch.manual_seed(42)
torch.set_num_threads(1)
mixing = torch.randn(3, 12)
data = torch.randn(512, 3) @ mixing + 0.05 * torch.randn(512, 12)
mean = data[:384].mean(0)
scale = data[:384].std(0).clamp_min(1e-4)
data = (data - mean) / scale
model = nn.Sequential(nn.Linear(24, 64), nn.GELU(), nn.Linear(64, 64),
                      nn.GELU(), nn.Linear(64, 12))
opt = torch.optim.Adam(model.parameters(), lr=0.005)

def make_mask(n):
    # Exactly six hidden coordinates per example avoids empty-loss batches.
    order = torch.rand(n, 12).argsort(1)
    hidden = torch.zeros(n, 12, dtype=torch.bool)
    return hidden.scatter(1, order[:, :6], True)

def reconstruct(values, hidden):
    visible = values.masked_fill(hidden, 0)
    return model(torch.cat((visible, hidden.float()), dim=-1))

fixed_mask = make_mask(128)
for _ in range(220):
    ids = torch.randint(384, (96,))
    values = data[ids]
    hidden = make_mask(len(values))
    opt.zero_grad(set_to_none=True)
    prediction = reconstruct(values, hidden)
    loss = (prediction - values).square()[hidden].mean()
    loss.backward()
    opt.step()
model.eval()
with torch.no_grad():
    test = data[384:]
    prediction = reconstruct(test, fixed_mask)
    mse = (prediction - test).square()[fixed_mask].mean().item()
    baseline = test.square()[fixed_mask].mean().item()
    corrupted = test.clone()
    corrupted[fixed_mask] = 999
    torch.testing.assert_close(reconstruct(corrupted, fixed_mask), prediction)
assert mse < baseline * 0.7
print({"masked_test_mse": mse, "train_mean_baseline_mse": baseline})
