Skip to content

Instantly share code, notes, and snippets.

@proger
Last active December 7, 2023 13:37
Show Gist options
  • Select an option

  • Save proger/f1bb94e3d2d9209a9024ba58821b24c0 to your computer and use it in GitHub Desktop.

Select an option

Save proger/f1bb94e3d2d9209a9024ba58821b24c0 to your computer and use it in GitHub Desktop.
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
import gzip
import numpy as np
torch.set_float32_matmul_precision('high')
def read(filename):
with gzip.open(filename, 'rb') as file:
compressed_data = file.read()
data = np.frombuffer(compressed_data, dtype=np.float32)
print(filename)
return data.reshape(-1, 100)
class MinibatchKMeans(nn.Module):
def __init__(self, k, d=100):
super().__init__()
self.k = k
self.d = d
self.centroids = nn.Parameter(torch.randn(k, d), requires_grad=False)
self.centroids_count = nn.Parameter(torch.zeros(k, dtype=torch.long), requires_grad=False)
@torch.no_grad()
def forward(self, x):
distances = torch.cdist(x, self.centroids)
min = torch.min(distances, dim=1)
#print(min)
if self.training:
for i in range(len(x)):
count = self.centroids_count[min.indices[i]]
self.centroids_count[min.indices[i]] = count+1
lr = 1/(count+1)
self.centroids[min.indices[i]] = (1-lr)*self.centroids[min.indices[i]] + lr*x[i]
return torch.mean(min.values)
filenames = """
./exp/embed/1/data\_split\_ad/stdout
./exp/embed/1/data\_split\_an/stdout
./exp/embed/1/data\_split\_az/stdout
./exp/embed/1/data\_split\_ap/stdout
./exp/embed/1/data\_split\_bf/stdout
./exp/embed/1/data\_split\_ba/stdout
./exp/embed/1/data\_split\_aw/stdout
./exp/embed/1/data\_split\_ac/stdout
./exp/embed/1/data\_split\_ai/stdout
./exp/embed/1/data\_split\_as/stdout
./exp/embed/1/data\_split\_ay/stdout
./exp/embed/1/data\_split\_am/stdout
./exp/embed/1/data\_split\_ag/stdout
./exp/embed/1/data\_split\_be/stdout
./exp/embed/1/data\_split\_bb/stdout
./exp/embed/1/data\_split\_aj/stdout
./exp/embed/1/data\_split\_at/stdout
./exp/embed/1/data\_split\_av/stdout
./exp/embed/1/data\_split\_ah/stdout
./exp/embed/1/data\_split\_ab/stdout
./exp/embed/1/data\_split\_ao/stdout
./exp/embed/1/data\_split\_ae/stdout
./exp/embed/1/data\_split\_aq/stdout
./exp/embed/1/data\_split\_bc/stdout
./exp/embed/1/data\_split\_aa/stdout
./exp/embed/1/data\_split\_ak/stdout
./exp/embed/1/data\_split\_au/stdout
./exp/embed/1/data\_split\_ax/stdout
./exp/embed/1/data\_split\_ar/stdout
./exp/embed/1/data\_split\_af/stdout
./exp/embed/1/data\_split\_al/stdout
./exp/embed/1/data\_split\_bd/stdout
"""
arrays = np.concatenate([read(name) for name in filenames.strip().split()[:]])
dataset = TensorDataset(torch.from_numpy(arrays))
k = 2**16
init_loader = DataLoader(dataset, batch_size=k, shuffle=True)
train_loader = DataLoader(dataset, batch_size=1024, shuffle=True)
device = 'cuda:1'
kmeans = MinibatchKMeans(k).to(device)
kmeans.centroids.data = next(iter(init_loader))[0].to(device)
kmeans = torch.compile(kmeans)
print('compile')
kmeans.train()
for i, batch in enumerate(train_loader):
x = batch[0].to(device)
loss = kmeans(x)
print(f'step {i}/{len(train_loader)} loss {loss.item()}')
torch.save(kmeans.state_dict(), 'exp/kmeans.pt')
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment