Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions configs/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ defaults:
- model: unet
- scheduler: linear
- solver: heun
- time_sampler: uniform
- optimizer: adam
- lr_scheduler: none
- dataset: mnist
Expand Down
3 changes: 3 additions & 0 deletions configs/time_sampler/logit_normal.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
type: logit_normal
mean: 0.0
stddev: 1.0
3 changes: 3 additions & 0 deletions configs/time_sampler/uniform.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
type: uniform
low: 0.0
high: 1.0
39 changes: 31 additions & 8 deletions scripts/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
import matplotlib.pyplot as plt
import mlflow
from _common import create_lr_scheduler, create_scheduler, create_solver, get_preprocess_hook
from loguru import logger
from omegaconf import DictConfig, OmegaConf
from sklearn.datasets import make_moons
from tinygrad.nn.optim import Adam
Expand All @@ -42,6 +43,7 @@
from tinyflow.losses import mse
from tinyflow.nn import NeuralNetwork, UNetCIFAR10, UNetCIFAR10Large, UNetMNIST
from tinyflow.path import AffinePath
from tinyflow.time_sampler import BaseTimeSampler, LogitNormalSampler, UniformTimeSampler
from tinyflow.trainer import CIFAR10Trainer, MNISTTrainer
from tinyflow.utils import visualize_moons

Expand Down Expand Up @@ -108,7 +110,23 @@ def create_dataloader(cfg: DictConfig):
raise ValueError(f"Unknown dataset type: {dataset_type}")


def create_trainer(cfg: DictConfig, model, dataloader, optim, path, lr_scheduler=None):
def create_time_sampler(cfg: DictConfig) -> BaseTimeSampler:
"""Create time sampler instance"""
time_sampler_type = cfg.time_sampler.get("type")
if time_sampler_type == "uniform":
low: float = cfg.time_sampler.get("low")
high: float = cfg.time_sampler.get("high")
return UniformTimeSampler(low, high)
if time_sampler_type == "logit_normal":
mean: float = cfg.time_sampler.get("mean")
stddev: float = cfg.time_sampler.get("stddev")
return LogitNormalSampler(mean, stddev)
raise ValueError(f"Unknown time sampler: {time_sampler_type}")


def create_trainer(
cfg: DictConfig, model, time_sampler, dataloader, optim, path, lr_scheduler=None
):
"""Create trainer from config (image datasets only)."""
dataset_type = cfg.dataset.get("type", cfg.dataset.name)
num_epochs = cfg.training.num_epochs
Expand All @@ -131,6 +149,7 @@ def create_trainer(cfg: DictConfig, model, dataloader, optim, path, lr_scheduler
num_epochs=num_epochs,
log_interval=log_interval,
lr_scheduler=lr_scheduler,
time_sampler=time_sampler,
gradient_accumulation_steps=gradient_accumulation_steps,
)

Expand All @@ -140,7 +159,8 @@ def generate_model_name(cfg: DictConfig) -> str:
dataset = cfg.dataset.get("type", cfg.dataset.name)
model_type = cfg.model.type
scheduler = cfg.scheduler.type.replace("Scheduler", "").lower()
return f"model_{dataset}_{model_type}_{scheduler}.safetensors"
time_sampler = cfg.time_sampler.type.lower()
return f"model_{dataset}_{model_type}_ts_{time_sampler}_{scheduler}.safetensors"


def train_moons(cfg: DictConfig, model_name: str):
Expand All @@ -152,6 +172,7 @@ def train_moons(cfg: DictConfig, model_name: str):
optim = Adam(get_parameters(model), lr=cfg.optimizer.lr)
lr_scheduler = create_lr_scheduler(cfg, optim)

time_sampler = create_time_sampler(cfg)
loss_fn = mse
losses = []

Expand All @@ -160,7 +181,7 @@ def train_moons(cfg: DictConfig, model_name: str):
for iter_idx in pbar:
x, _ = make_moons(n_samples=cfg.dataset.n_samples, noise=cfg.dataset.noise)
x_1 = T(x.astype("float32")) # pyright: ignore
t = T.rand(x_1.shape[0], 1) * 0.99 # clamping
t = time_sampler.sample(x_1.shape[0], 1) # T.rand(x_1.shape[0], 1) * 0.99 # clamping
x_0 = T.randn(*x_1.shape)
x_t, dx_t = path.sample(x_1=x_1, t=t, x_0=x_0)
out = model(x_t, t=t) # pyright: ignore
Expand All @@ -184,7 +205,7 @@ def train_moons(cfg: DictConfig, model_name: str):

if cfg.training.get("save_model", True):
safe_save(get_state_dict(model), model_name)
print(f"✓ Model saved to: {model_name}")
logger.info(f"✓ Model saved to: {model_name}")

if cfg.training.get("log_artifacts", True):
output_dir = cfg.get("output_dir", "outputs")
Expand Down Expand Up @@ -226,7 +247,8 @@ def train_images(cfg: DictConfig, model_name: str):
lr_scheduler = create_lr_scheduler(cfg, optim)

dataloader = create_dataloader(cfg)
trainer = create_trainer(cfg, model, dataloader, optim, path, lr_scheduler)
time_sampler = create_time_sampler(cfg)
trainer = create_trainer(cfg, model, time_sampler, dataloader, optim, path, lr_scheduler)

dataset_name = cfg.dataset.get("type", cfg.dataset.name)
with mlflow.start_run(run_name=dataset_name):
Expand All @@ -248,14 +270,15 @@ def train_images(cfg: DictConfig, model_name: str):
@hydra.main(version_base=None, config_path="../configs", config_name="config")
def main(cfg: DictConfig):
"""Hydra entry point; dispatches on dataset type."""
print("Configuration:")
print(OmegaConf.to_yaml(cfg))
logger.info("Configuration:")
logger.info(OmegaConf.to_yaml(cfg))

if cfg.get("seed"):
T.manual_seed(cfg.seed)

model_name = cfg.get("model_name", generate_model_name(cfg))
print(f"\nModel will be saved as: {model_name}")
model_name = f"weights/{model_name}"
logger.info(f"\nModel will be saved as: {model_name}")

dataset_type = cfg.dataset.get("type", cfg.dataset.name)
if dataset_type == "moons":
Expand Down
101 changes: 101 additions & 0 deletions tests/test_time_sampler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,101 @@
import numpy as np
import pytest
from tinygrad.tensor import Tensor as T

from tinyflow.time_sampler import BaseTimeSampler, LogitNormalSampler, UniformTimeSampler


class TestBaseTimeSampler:
def test_sample_not_implemented(self):
"""Base class sample() must be overridden by subclasses"""
sampler = BaseTimeSampler()
with pytest.raises(NotImplementedError):
sampler.sample(4, 1)


class TestUniformTimeSampler:
def test_default_output_shape(self):
"""Sampled tensor shape matches requested shape"""
sampler = UniformTimeSampler()
t = sampler.sample(16, 1)

assert t.shape == (16, 1)

def test_default_bounds(self):
"""Default low=0.0, high=1.0 bounds all samples in [0, 1)"""
T.manual_seed(0)
sampler = UniformTimeSampler()
t = sampler.sample(1000, 1).numpy()

assert t.min() >= 0.0
assert t.max() < 1.0

def test_custom_bounds(self):
"""Custom low/high bounds are respected"""
T.manual_seed(0)
sampler = UniformTimeSampler(low=0.2, high=0.7)
t = sampler.sample(1000, 1).numpy()

assert t.min() >= 0.2
assert t.max() < 0.7

def test_legacy_clamp_range(self):
"""low=0.0, high=0.99 reproduces the historical `T.rand(...) * 0.99` clamp"""
T.manual_seed(0)
sampler = UniformTimeSampler(low=0.0, high=0.99)
t = sampler.sample(1000, 1).numpy()

assert t.min() >= 0.0
assert t.max() < 0.99


class TestLogitNormalSampler:
def test_output_shape(self):
"""Sampled tensor shape matches requested shape"""
sampler = LogitNormalSampler()
t = sampler.sample(16, 1)

assert t.shape == (16, 1)

def test_output_in_unit_interval(self):
"""sigmoid(.) squashes all samples into the open interval (0, 1)"""
T.manual_seed(0)
sampler = LogitNormalSampler(mean=0.0, stddev=1.0)
t = sampler.sample(1000, 1).numpy()

assert t.min() > 0.0
assert t.max() < 1.0

def test_default_mean_is_centered(self):
"""m=0.0 is the median of the underlying normal, so sigmoid(0)=0.5 is
the median of the sampled distribution"""
T.manual_seed(0)
sampler = LogitNormalSampler(mean=0.0, stddev=1.0)
t = sampler.sample(5000, 1).numpy()

median = np.median(t)
assert abs(median - 0.5) < 0.05

def test_shifted_mean_skews_distribution(self):
"""A positive m shifts probability mass toward t=1"""
T.manual_seed(0)
low_m_sampler = LogitNormalSampler(mean=-2.0, stddev=1.0)
high_m_sampler = LogitNormalSampler(mean=2.0, stddev=1.0)

t_low = low_m_sampler.sample(2000, 1).numpy()
t_high = high_m_sampler.sample(2000, 1).numpy()

assert t_low.mean() < 0.5 < t_high.mean()

def test_small_scale_concentrates_near_median(self):
"""A small s concentrates samples tightly around sigmoid(m)"""
T.manual_seed(0)
sampler = LogitNormalSampler(mean=0.0, stddev=0.01)
t = sampler.sample(1000, 1).numpy()

assert abs(t.mean() - 0.5) < 0.05
assert t.std() < 0.05


if __name__ == "__main__":
pytest.main([__file__, "-v"])
1 change: 0 additions & 1 deletion tinyflow/solver/rk4.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@ def __init__(

@logger.catch(reraise=True)
def sample(self, h, t, rhs_prev):
# t = self.preprocess_hook(t, rhs_prev)
return self.step(h, t, rhs_prev)

@logger.catch(reraise=True)
Expand Down
34 changes: 34 additions & 0 deletions tinyflow/time_sampler.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,34 @@
from loguru import logger
from tinygrad.tensor import Tensor


class BaseTimeSampler:
def __init__(self):
pass

def sample(self, *shape):
"""Perform sampling operation"""
raise NotImplementedError


class UniformTimeSampler(BaseTimeSampler):
def __init__(self, low: float = 0.0, high: float = 1.0):
self.l = low
self.h = high

@logger.catch(reraise=True)
def sample(self, *shape) -> Tensor:
out = Tensor.rand(*shape)
return (self.h - self.l) * out + self.l


class LogitNormalSampler(BaseTimeSampler):
def __init__(self, mean: float = 0.0, stddev: float = 1.0):
self.m = mean
self.s = stddev

@logger.catch(reraise=True)
def sample(self, *shape) -> Tensor:
out: Tensor = self.m + self.s * Tensor.randn(*shape)
out: Tensor = out.sigmoid()
return out
9 changes: 7 additions & 2 deletions tinyflow/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from tinyflow.nn import BaseNeuralNetwork
from tinyflow.path import Path
from tinyflow.solver import ODESolver
from tinyflow.time_sampler import BaseTimeSampler
from tinyflow.utils import visualize_cifar10, visualize_mnist

# Import metrics lazily to avoid circular imports
Expand All @@ -33,6 +34,7 @@ def __init__(
num_epochs: int = 10_000,
log_interval: int = 50,
lr_scheduler=None,
time_sampler: BaseTimeSampler | None = None,
gradient_accumulation_steps: int = 1,
compute_fid: bool = False,
fid_interval: int = 50,
Expand All @@ -48,6 +50,9 @@ def __init__(
self.num_epochs = num_epochs
self.log_interval = log_interval
self.lr_scheduler = lr_scheduler
if time_sampler is None:
raise ValueError("Time sampler cannot be None")
self.time_sampler = time_sampler
self.gradient_accumulation_steps = gradient_accumulation_steps
self._losses = []
self.global_step = 0
Expand Down Expand Up @@ -216,7 +221,7 @@ def epoch(self, epoch_idx: int | None):
for batch in tqdm(self.dataloader, desc=desc):
x_batch, _ = batch
x = T(x_batch)
t = T.rand(x.shape[0], 1) * 0.99
t = self.time_sampler.sample(x.shape[0], 1)
x_0 = T.randn(*x.shape)
x_t, dx_t = self.path.sample(x_1=x, t=t, x_0=x_0)
out = self.model(x_t, t)
Expand Down Expand Up @@ -286,7 +291,7 @@ def epoch(self, epoch_idx: int | None):
x_batch, _ = batch
x = T(x_batch) # Already float32 from dataloader

t = T.rand(x.shape[0], 1) * 0.99
t = self.time_sampler.sample(x.shape[0], 1)
x_0 = T.randn(*x.shape)

x_t, dx_t = self.path.sample(x_1=x, t=t, x_0=x_0)
Expand Down
Loading