From 59163ec32d09a7fc8f75f7eefe0f1634db4e2440 Mon Sep 17 00:00:00 2001 From: Ilia Ilmer Date: Sun, 26 Jul 2026 12:22:22 -0400 Subject: [PATCH 1/4] add time sampler --- configs/config.yaml | 1 + configs/time_sampler/logit_normal.yaml | 3 +++ configs/time_sampler/uniform.yaml | 3 +++ scripts/train.py | 35 ++++++++++++++++++++------ tinyflow/time_sampler.py | 34 +++++++++++++++++++++++++ tinyflow/trainer.py | 9 +++++-- 6 files changed, 76 insertions(+), 9 deletions(-) create mode 100644 configs/time_sampler/logit_normal.yaml create mode 100644 configs/time_sampler/uniform.yaml create mode 100644 tinyflow/time_sampler.py diff --git a/configs/config.yaml b/configs/config.yaml index fe06137..e1d2af0 100644 --- a/configs/config.yaml +++ b/configs/config.yaml @@ -4,6 +4,7 @@ defaults: - model: unet - scheduler: linear - solver: heun + - time_sampler: uniform - optimizer: adam - lr_scheduler: none - dataset: mnist diff --git a/configs/time_sampler/logit_normal.yaml b/configs/time_sampler/logit_normal.yaml new file mode 100644 index 0000000..0696511 --- /dev/null +++ b/configs/time_sampler/logit_normal.yaml @@ -0,0 +1,3 @@ +type: logit_normal +low: 0.0 +high: 1.0 diff --git a/configs/time_sampler/uniform.yaml b/configs/time_sampler/uniform.yaml new file mode 100644 index 0000000..b09d1f8 --- /dev/null +++ b/configs/time_sampler/uniform.yaml @@ -0,0 +1,3 @@ +type: uniform +low: 0.0 +high: 1.0 diff --git a/scripts/train.py b/scripts/train.py index 150c933..d784177 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -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 @@ -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 @@ -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": + low: float = cfg.time_sampler.get("low") + high: float = cfg.time_sampler.get("high") + return LogitNormalSampler(low, high) + 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 @@ -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, ) @@ -152,6 +171,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 = [] @@ -160,7 +180,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 @@ -184,7 +204,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") @@ -226,7 +246,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): @@ -248,14 +269,14 @@ 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}") + logger.info(f"\nModel will be saved as: {model_name}") dataset_type = cfg.dataset.get("type", cfg.dataset.name) if dataset_type == "moons": diff --git a/tinyflow/time_sampler.py b/tinyflow/time_sampler.py new file mode 100644 index 0000000..78f8ae2 --- /dev/null +++ b/tinyflow/time_sampler.py @@ -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, m: float = 0.0, s: float = 1.0): + self.m = m + self.s = s + + @logger.catch(reraise=True) + def sample(self, *shape) -> Tensor: + out: Tensor = self.m + self.s * Tensor.randn(*shape) + out: Tensor = out.sigmoid() + return out diff --git a/tinyflow/trainer.py b/tinyflow/trainer.py index b7843d6..ea55d23 100644 --- a/tinyflow/trainer.py +++ b/tinyflow/trainer.py @@ -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 @@ -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, @@ -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 @@ -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) @@ -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) From 17453ab76c22f841951fb1787b2ee726e47eb655 Mon Sep 17 00:00:00 2001 From: Ilia Ilmer Date: Sun, 26 Jul 2026 12:26:11 -0400 Subject: [PATCH 2/4] rename+bugifx --- configs/time_sampler/logit_normal.yaml | 4 ++-- scripts/train.py | 10 ++++++---- tinyflow/time_sampler.py | 6 +++--- 3 files changed, 11 insertions(+), 9 deletions(-) diff --git a/configs/time_sampler/logit_normal.yaml b/configs/time_sampler/logit_normal.yaml index 0696511..d2356f8 100644 --- a/configs/time_sampler/logit_normal.yaml +++ b/configs/time_sampler/logit_normal.yaml @@ -1,3 +1,3 @@ type: logit_normal -low: 0.0 -high: 1.0 +mean: 0.0 +stddev: 1.0 diff --git a/scripts/train.py b/scripts/train.py index d784177..6901b17 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -118,9 +118,9 @@ def create_time_sampler(cfg: DictConfig) -> BaseTimeSampler: high: float = cfg.time_sampler.get("high") return UniformTimeSampler(low, high) if time_sampler_type == "logit_normal": - low: float = cfg.time_sampler.get("low") - high: float = cfg.time_sampler.get("high") - return LogitNormalSampler(low, high) + 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}") @@ -159,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): @@ -276,6 +277,7 @@ def main(cfg: DictConfig): T.manual_seed(cfg.seed) model_name = cfg.get("model_name", generate_model_name(cfg)) + 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) diff --git a/tinyflow/time_sampler.py b/tinyflow/time_sampler.py index 78f8ae2..bf84136 100644 --- a/tinyflow/time_sampler.py +++ b/tinyflow/time_sampler.py @@ -23,9 +23,9 @@ def sample(self, *shape) -> Tensor: class LogitNormalSampler(BaseTimeSampler): - def __init__(self, m: float = 0.0, s: float = 1.0): - self.m = m - self.s = s + 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: From a067866ef9eb2fcfe556832b24f0d9cc752dfba8 Mon Sep 17 00:00:00 2001 From: Ilia Ilmer Date: Sun, 26 Jul 2026 12:29:01 -0400 Subject: [PATCH 3/4] add tests --- tests/test_time_sampler.py | 101 +++++++++++++++++++++++++++++++++++++ 1 file changed, 101 insertions(+) create mode 100644 tests/test_time_sampler.py diff --git a/tests/test_time_sampler.py b/tests/test_time_sampler.py new file mode 100644 index 0000000..a90eb7d --- /dev/null +++ b/tests/test_time_sampler.py @@ -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"]) From 32ab05f3c9e11fe94f7f4418c583099fe591bcc4 Mon Sep 17 00:00:00 2001 From: Ilia Ilmer Date: Sun, 26 Jul 2026 12:30:44 -0400 Subject: [PATCH 4/4] remove dead code --- tinyflow/solver/rk4.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tinyflow/solver/rk4.py b/tinyflow/solver/rk4.py index 07682dc..5e15389 100644 --- a/tinyflow/solver/rk4.py +++ b/tinyflow/solver/rk4.py @@ -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)