Browse Source

[MNT] Faster Training

tags/v0.3.2
shihy 2 years ago
parent
commit
1f108d88de
4 changed files with 118 additions and 87 deletions
  1. +25
    -15
      examples/dataset_cifar_workflow/benchmarks/dataset/data.py
  2. +6
    -1
      examples/dataset_cifar_workflow/benchmarks/dataset/utils.py
  3. +76
    -65
      examples/dataset_cifar_workflow/benchmarks/utils.py
  4. +11
    -6
      learnware/specification/regular/image/rkme.py

+ 25
- 15
examples/dataset_cifar_workflow/benchmarks/dataset/data.py View File

@@ -5,27 +5,37 @@ import torch
from torch.utils.data import random_split, Subset
from torchvision import datasets
from torchvision.transforms import transforms
from torch.utils.data import TensorDataset

from .utils import cached
from examples.dataset_cifar_workflow.benchmarks.dataset.utils import split_dataset, build_transforms

cache_root = os.path.abspath(os.path.join(os.path.dirname( __file__ ), '..', '..', 'cache'))

cifar_data = torch.stack([u[0] for u in datasets.CIFAR10(root="cache", download=True,
train=True, transform=transforms.ToTensor())])
augment_transform, regular_transform, whiten_transform = build_transforms(cifar_data)

cifar_train_set_augment = datasets.CIFAR10(root="cache", download=True,
train=True, transform=whiten_transform)
cifar_test_set = datasets.CIFAR10(root="cache", download=True,
train=False, transform=whiten_transform)
cifar_spec_train_set = datasets.CIFAR10(root="cache", download=True,
train=True, transform=whiten_transform)
cifar_spec_test_set = datasets.CIFAR10(root="cache", download=True,
train=False, transform=whiten_transform)
cifar_train = datasets.CIFAR10(root=cache_root, download=True, train=True, transform=transforms.ToTensor())
cifar_train_X = torch.stack([u[0] for u in cifar_train])
augment_transform, regular_transform, whiten_transform = build_transforms(cifar_train_X)

cifar_train_set_augment = datasets.CIFAR10(root=cache_root, download=True, train=True, transform=whiten_transform)
cifar_test_set = datasets.CIFAR10(root=cache_root, download=True, train=False, transform=whiten_transform)
cifar_spec_train_set = datasets.CIFAR10(root=cache_root, download=True, train=True, transform=whiten_transform)
cifar_spec_test_set = datasets.CIFAR10(root=cache_root, download=True, train=False, transform=whiten_transform)
train_targets = cifar_train_set_augment.targets
test_targets = cifar_test_set.targets

def faster_train(device):
global cifar_train_set_augment
global cifar_test_set
global cifar_spec_train_set
global cifar_spec_test_set
cifar_train_set_augment = cached(cifar_train_set_augment, device=device)
cifar_test_set = cached(cifar_test_set, device=device)
cifar_spec_train_set = cached(cifar_spec_train_set, device=device)
cifar_spec_test_set = cached(cifar_spec_test_set, device=device)

def uploader_data(order=None):
train_indices, order = split_dataset(torch.asarray(cifar_train_set_augment.targets), 12500, split="uploader", order=order)
valid_indices, _ = split_dataset(torch.asarray(cifar_test_set.targets), 2000, split="uploader", order=order)
train_indices, order = split_dataset(torch.asarray(train_targets), 12500, split="uploader", order=order)
valid_indices, _ = split_dataset(torch.asarray(test_targets), 2000, split="uploader", order=order)

return (Subset(cifar_train_set_augment, train_indices),
Subset(cifar_test_set, valid_indices),
@@ -34,6 +44,6 @@ def uploader_data(order=None):

def user_data(indices=None, order=None):
if indices is None:
indices, order = split_dataset(torch.asarray(cifar_spec_test_set.targets), 3000, split="user", order=order)
indices, order = split_dataset(torch.asarray(test_targets), 3000, split="user", order=order)

return Subset(cifar_test_set, indices), Subset(cifar_spec_test_set, indices), indices, order

+ 6
- 1
examples/dataset_cifar_workflow/benchmarks/dataset/utils.py View File

@@ -4,6 +4,7 @@ from functools import reduce
import numpy as np
import torch
import torchvision
from torch.utils.data import TensorDataset, Dataset, DataLoader

torchvision.disable_beta_transforms_warning()
from torchvision.transforms import transforms, v2
@@ -85,4 +86,8 @@ def build_transforms(train_X):
transforms.LinearTransformation(whitening_matrix, torch.zeros_like(train_X[0].reshape(-1)))
])

return augment_transform, regular_transform, whiten_transform
return augment_transform, regular_transform, whiten_transform

def cached(data: Dataset, device):
X, y = next(iter(DataLoader(data, batch_size=len(data))))
return TensorDataset(X.to(device), y.to(device))

+ 76
- 65
examples/dataset_cifar_workflow/benchmarks/utils.py View File

@@ -14,12 +14,15 @@ from learnware.client import LearnwareClient
from learnware.learnware import Learnware
from learnware.specification import generate_rkme_image_spec, RKMEImageSpecification
from .dataset import uploader_data, user_data
from .dataset.utils import cached
from .models.conv import ConvModel
from learnware.market import LearnwareMarket
from learnware.utils import choose_device

from torch.profiler import profile, record_function, ProfilerActivity

@torch.no_grad()
def evaluate(model, evaluate_set: Dataset, device=None):
def evaluate(model, evaluate_set: Dataset, device=None, distribution=True):
device = choose_device(0) if device is None else device

if isinstance(model, nn.Module):
@@ -29,16 +32,20 @@ def evaluate(model, evaluate_set: Dataset, device=None):
mapping = lambda m, x: m.predict(x)

criterion = nn.CrossEntropyLoss(reduction="sum")
total, correct, loss = 0, 0, 0.0
dataloader = DataLoader(evaluate_set, batch_size=512, shuffle=True)
total, correct, loss = 0, 0, torch.as_tensor(0.0, dtype=torch.float32, device=device)
dataloader = DataLoader(evaluate_set, batch_size=1024, shuffle=True)
for i, (X, y) in enumerate(dataloader):
X, y = X.to(device), y.to(device)
out = mapping(model, X)
if not torch.is_tensor(out):
out = torch.from_numpy(out).to(device)
loss += criterion(out, y)

_, predicted = torch.max(out.data, 1)
if distribution:
loss += criterion(out, y)
_, predicted = torch.max(out.data, 1)
else:
predicted = out

total += y.size(0)
correct += (predicted == y).sum().item()

@@ -67,56 +74,17 @@ def build_learnware(name: str, market: LearnwareMarket, order, model_name="conv"

channel = train_set[0][0].shape[0]
image_size = train_set[0][0].shape[1], train_set[0][0].shape[2]

model = ConvModel(channel=channel, im_size=image_size,
n_random_features=out_classes).to(device)

model.train()

# SGD optimizer with learning rate 1e-2
optimizer = optim.SGD(model.parameters(), lr=1e-2, momentum=0.9)
# Scheduler
# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20)
# mean-squared error loss
criterion = nn.CrossEntropyLoss()
# Prepare DataLoader
dataloader = DataLoader(train_set, batch_size=batch_size, shuffle=True)
# valid loss
best_loss = 100000 # initially
# Optimizing...
for epoch in range(epochs):
running_loss = []
model.train()
for i, (X, y) in enumerate(dataloader):
X, y = X.to(device=device), y.to(device=device)
optimizer.zero_grad()
out = model(X)
loss = criterion(out, y)
loss.backward()
optimizer.step()
running_loss.append(loss.item())

valid_loss, valid_acc = evaluate(model, valid_set, device=device)
train_loss, train_acc = evaluate(model, train_set, device=device)
if valid_loss < best_loss:
best_loss = valid_loss

torch.save(model.state_dict(), os.path.join(cache_dir, "model.pth"))
print("Epoch: {}, Valid Best Accuracy: {:.3f}% ({:.3f})".format(epoch+1, valid_acc, valid_loss))
if valid_acc > 99.0:
print("Early Stopping at 99% !")
break

if (epoch + 1) % 5 == 0:
print('Epoch: {}, Train Average Loss: {:.3f}, Accuracy {:.3f}%, Valid Average Loss: {:.3f}'.format(
epoch+1, np.mean(running_loss), train_acc, valid_loss))

# scheduler.step()
# train model
save_path = os.path.join(cache_dir, "model.pth")
train_model(model, train_set, valid_set, save_path, epochs=epochs, batch_size=batch_size, device=device)

# build specification
loader = DataLoader(spec_set, batch_size=3000, shuffle=True)
sampled_X, _ = next(iter(loader))
spec = generate_rkme_image_spec(sampled_X, whitening=False, cross_platform=False)
spec = generate_rkme_image_spec(sampled_X, whitening=False, experimental=True)

# add to market
model_dir = os.path.abspath(os.path.join(__file__, "..", "models"))
@@ -158,6 +126,49 @@ def build_learnware(name: str, market: LearnwareMarket, order, model_name="conv"

return model

def train_model(model: nn.Module, train_set: Dataset, valid_set: Dataset,
save_path: str, epochs=35, batch_size=128, device=None):
device = choose_device(0) if device is None else device

model.train()
# SGD optimizer with learning rate 1e-2
optimizer = optim.SGD(model.parameters(), lr=1e-2, momentum=0.9)
# Scheduler
# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20)
# mean-squared error loss
criterion = nn.CrossEntropyLoss()
# Prepare DataLoader
dataloader = DataLoader(train_set, batch_size=batch_size, shuffle=True)
# valid loss
best_loss = 100000 # initially
# Optimizing...
for epoch in range(epochs):
running_loss = []
model.train()
for i, (X, y) in enumerate(dataloader):
X, y = X.to(device=device), y.to(device=device)
optimizer.zero_grad()
out = model(X)
loss = criterion(out, y)
loss.backward()
optimizer.step()
running_loss.append(loss.item())

valid_loss, valid_acc = evaluate(model, valid_set, device=device)
train_loss, train_acc = evaluate(model, train_set, device=device)
if valid_loss < best_loss:
best_loss = valid_loss

torch.save(model.state_dict(), save_path)
print("Epoch: {}, Valid Best Accuracy: {:.3f}% ({:.3f})".format(epoch+1, valid_acc, valid_loss))
if valid_acc > 99.0:
print("Early Stopping at 99% !")
break

if (epoch + 1) % 5 == 0:
print('Epoch: {}, Train Average Loss: {:.3f}, Accuracy {:.3f}%, Valid Average Loss: {:.3f}'.format(
epoch+1, np.mean(running_loss), train_acc, valid_loss))


def build_specification(name: str, cache_id, order, sampled_size=3000):
cache_dir = os.path.abspath(os.path.join(
@@ -174,7 +185,7 @@ def build_specification(name: str, cache_id, order, sampled_size=3000):
test_dataset, spec_dataset, indices, _ = user_data(order=order)
loader = DataLoader(spec_dataset, batch_size=sampled_size, shuffle=True)
sampled_X, _ = next(iter(loader))
spec = generate_rkme_image_spec(sampled_X, whitening=False, cross_platform=False)
spec = generate_rkme_image_spec(sampled_X, whitening=False, experimental=True)

spec.msg = indices.tolist()
spec.save(cache_path)
@@ -184,20 +195,14 @@ def build_specification(name: str, cache_id, order, sampled_size=3000):

class Recorder:

def __init__(self):
def __init__(self, headers, formats):
assert len(headers) == len(formats)
self.data = defaultdict(list)
self.headers = headers
self.formats = formats

def record(self, name, accuracy, loss):
self.data[name].append((accuracy, loss))

def latest(self):
table = []

for name, values in self.data.items():
value = values[-1]
table.append([name, "{:.3f}%".format(value[0]), "{:.3f}".format(value[1])])

return str(tabulate(table, headers=["Case", "Accuracy", "Loss"], tablefmt='orgtbl'))
def record(self, name, *args):
self.data[name].append(args)

def summary(self):
table = []
@@ -205,8 +210,14 @@ class Recorder:
for name, values in self.data.items():
value_mean = [np.mean(v) for v in zip(*values)]
value_std = [np.std(v) for v in zip(*values)]
table.append([name,
"{:.3f}% ± {:.3f}%".format(value_mean[0], value_std[0]),
"{:.3f} ± {:.3f}" .format(value_mean[1], value_std[1])])
table.append([name] + [f.format(m, s) for f, m, s in zip(self.formats, value_mean, value_std)])

return str(tabulate(table, headers=["Case"] + self.headers, tablefmt='orgtbl'))

def save(self, path):
with open(path, "w") as f:
json.dump(self.data, f)

return str(tabulate(table, headers=["Case", "Accuracy", "Loss"], tablefmt='orgtbl'))
def load(self, path):
with open(path, "r") as f:
self.data = json.load(f)

+ 11
- 6
learnware/specification/regular/image/rkme.py View File

@@ -168,7 +168,8 @@ class RKMEImageSpecification(RegularStatSpecification):
raise ModuleNotFoundError(
f"RKMEImageSpecification is not available because 'torch-optimizer' is not installed! Please install it manually.")

cross_platform = "cross_platform" not in kwargs or kwargs["cross_platform"]
# Cross-platform by default, unless the spec is specified to be generated specifically for local experiments.
cross_platform = "experimental" not in kwargs or not kwargs["experimental"]
# crucial
with deterministic(cross_platform, self._device) as random_generator:
self._random_generator = random_generator
@@ -439,13 +440,17 @@ def _get_zca_matrix(X, reg_coef=0.1):

class RandomGenerator:

def __init__(self, seed=0):
def __init__(self, seed=0, cross_platform=True):
self.cross_platform=cross_platform
self.state = RandomState(seed)

def normal_(self, tensor: torch.Tensor, mean=0.0, std=1.0):
data = self.state.normal(mean, std, size=tensor.shape)
with torch.no_grad():
tensor.copy_(torch.asarray(data, dtype=tensor.dtype))
if self.cross_platform:
data = self.state.normal(mean, std, size=tensor.shape)
with torch.no_grad():
tensor.copy_(torch.asarray(data, dtype=tensor.dtype))
else:
torch.nn.init.normal_(tensor, mean, std)


@contextmanager
@@ -459,7 +464,7 @@ def deterministic(cross_platform, device):
new_state=torch.cuda.get_rng_state(device.index),
device="cpu")

yield RandomGenerator(0)
yield RandomGenerator(seed=0, cross_platform=cross_platform)

torch.backends.cudnn.deterministic = deterministic_state
if cross_platform:


Loading…
Cancel
Save