diff --git a/.gitignore b/.gitignore index b7faf40..eabff5a 100644 --- a/.gitignore +++ b/.gitignore @@ -205,3 +205,6 @@ cython_debug/ marimo/_static/ marimo/_lsp/ __marimo__/ + +# Data from examples +data/ \ No newline at end of file diff --git a/README.md b/README.md index df298bf..dcde0f8 100644 --- a/README.md +++ b/README.md @@ -1,25 +1,32 @@ -# EmotiGrad +# 🌈 EmotiGrad β€” Emotional Support for Your Optimizers + +

+ + + + + +

EmotiGrad is a tiny Python library that wraps your PyTorch optimizers and gives you emotionally-charged feedback during training, from wholesome encouragement to unhinged sass. It aims to be: -- **Drop-in friendly** – keep your usual `torch.optim` code -- **Fun but useful** – emotional logs + basic training insights -- **Extensible** – easily add new "personalities" and behaviors - -> ### Because sometimes you need more than just `.step()`, you need support. +* **Drop-in friendly**: swap it into any `torch.optim` workflow +* **Fun but useful**: emotional logs + basic training insights +* **Extensible**: easily add new "personalities" and behaviors ## Status -> ⚠️ EmotiGrad is under active early development (pre-release). -> The API may change before `0.1.0`. Feedback and ideas are very welcome! +> ⚠️ EmotiGrad is under active early development (pre-release). +> Expect the API to evolve before version `0.1.0`. +> Feedback and ideas are *very* welcome! --- ## Installation -For now, install from source: +Install from source: ```bash git clone git@github.com:smiley-maker/emotigrad.git @@ -27,70 +34,119 @@ cd emotigrad pip install -e . ``` -(PyPI support will come in a later release.) +PyPI packages will come in a later release. + +--- ## Quick Start -Basic Usage: +Here’s the smallest possible example: ```python import torch from emotigrad import EmotionalOptimizer model = torch.nn.Linear(10, 1) -base_optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) +base_opt = torch.optim.Adam(model.parameters(), lr=1e-3) # Wrap your optimizer with a personality -optimizer = EmotionalOptimizer( - base_optimizer, - personality="wholesome", # (planned) "sassy", "chaotic", etc. +opt = EmotionalOptimizer( + base_opt, + personality="wholesome", # also: "sassy", "quiet", custom callables, etc. + message_every=20, # feedback every 20 steps (averaged) ) -for step in range(10): +for step in range(50): x = torch.randn(32, 10) y = torch.randn(32, 1) preds = model(x) loss = (preds - y).pow(2).mean() - optimizer.optimizer.zero_grad() + opt.zero_grad() loss.backward() - # In a future version, this call will emit emotional messages - optimizer.step() + # Provide loss to trigger feedback + opt.step(loss=loss.item()) +``` + +### How `message_every` works + +Instead of reacting to every single (noisy) loss value, EmotiGrad: + +1. Collects the last **N** loss values +2. Computes the **average loss for that block** +3. Compares it to the **previous block’s average** +4. Feeds the result into your chosen personality + +This produces smoother, more meaningful emotional feedback. + +Set `message_every=1` for per-step chatter. + + +## Personalities + +EmotiGrad ships with several built-in personalities, such as: + +* **wholesome** – kind, encouraging, proud of your progress +* **sassy** – mildly offended by your gradients +* **quiet** – reports occasionally, like a stoic mentor + +You can also write your own: + +```python +def hype(loss, prev, step): + if prev and loss < prev: + return f"πŸš€ Step {step}: HUGE gains! {prev:.4f} β†’ {loss:.4f}" + return None + +opt = EmotionalOptimizer(base_opt, personality=hype) +``` + +Or register them globally: + +```python +from emotigrad.personalities import register_personality + +register_personality("hype", hype) +opt = EmotionalOptimizer(base_opt, personality="hype") ``` -In upcoming versions, `optimizer.step(loss=loss.item())` will trigger personality-specific messages based on how training is going. +## Examples -## Roadmap (high level) +You can find examples in the `examples/` directory: + +* `basic_usage.py` +* `mnist_training.py` +* `custom_personality.py` + +## Roadmap Planned features: -- Emotional personas: - - wholesome – positive, encouraging - - sassy – mildly offended by your gradients - - chaotic – unhelpful but entertaining -- Basic training trend detection (loss going up/down, plateauing) -- Configurable verbosity and logging destinations -- Easy hooks for custom personalities -- Longer-term ideas: - - Integration with PyTorch Lightning / HuggingFace Trainer - - Optional LLM-based "training advisor" for suggestions +* More built-in emotional personas: + * `wholesome`, `sassy`, `quiet`, `chaotic`, `roaster`, `nervous`, etc. +* Trend-aware training feedback +* Configurable output formatting (e.g. text colors and formatting) +* Easy hooks for custom personalities +* Optional LLM-based β€œtraining advisor” mode +* Integrations with: + * PyTorch Lightning + * HuggingFace Trainer +* Rich visual outputs (ASCII art, emoji graphs, etc.) -## Contributions +## Contributing -Contributions are very welcome, even at this early stage! +Contributions are warmly welcome β€” even small improvements help! -Some helpful ways to contribute: +Ways to contribute: -- Try EmotiGrad in a toy project and open issues for: - - bugs - - confusing APIs - - personality ideas - - Add or improve a personality preset -- Add tests or docs +* Report bugs or confusing APIs +* Suggest new personalities or features +* Improve tests or documentation +* Add examples -To setup for development: +To set up development: ```bash git clone git@github.com:smiley-maker/emotigrad.git @@ -99,23 +155,33 @@ pip install -e ".[dev]" pytest ``` -We also suggest using a virtual environment or conda to manage dependencies. Development dependencies and more detailed instructions will come as the project evolves. +We also recommend using a virtual environment or conda during local development. -## Current Project Structure + +## Project Structure ``` emotigrad/ - emotigrad/ - __init__.py - base.py # will hold EmotionalOptimizer + src/ + emotigrad/ + __init__.py + base.py # EmotionalOptimizer + personalities.py # built-in personas + registry + types.py # Personality Protocol tests/ - test_smoke.py # tiny test so CI has something + examples/ README.md LICENSE pyproject.toml - .gitignore ``` + ## License -EmotiGrad is open source and under the MIT License. Learn more in the LICENSE file. \ No newline at end of file +EmotiGrad is released under the MIT License. +See the LICENSE file for details. + + +## Thanks for checking out EmotiGrad + +If you build something with it, please share it or open an issue, I’d love to see what you make! diff --git a/examples/basic_usage.py b/examples/basic_usage.py new file mode 100644 index 0000000..dea7d23 --- /dev/null +++ b/examples/basic_usage.py @@ -0,0 +1,50 @@ +""" +Basic usage example for EmotiGrad. + +This script shows how to wrap a PyTorch optimizer with an EmotionalOptimizer +and get emotionally-enhanced training feedback. It uses a tiny synthetic dataset +so it runs instantly and without external dependencies. + +Run with: + python examples/basic_usage.py +""" + +import torch + +from emotigrad import EmotionalOptimizer + + +def main(): + # Simple linear model for demonstration + model = torch.nn.Linear(10, 1) + + # Standard optimizer + base_opt = torch.optim.Adam(model.parameters(), lr=1e-3) + + # Wrap with EmotiGrad! + opt = EmotionalOptimizer( + base_opt, + personality="wholesome", # try "sassy" or write your own! + message_every=5, # emotional feedback every 5 steps + ) + + # Synthetic training loop + for step in range(20): + x = torch.randn(32, 10) + y = torch.randn(32, 1) + + preds = model(x) + loss = (preds - y).pow(2).mean() + + opt.zero_grad() + loss.backward() + + # Passing loss triggers EmotiGrad's emotional feedback + opt.step(loss=loss.item()) + + if step % 5 == 0: + print(f"[step {step}] loss = {loss.item():.4f}") + + +if __name__ == "__main__": + main() diff --git a/examples/custom_personality.py b/examples/custom_personality.py new file mode 100644 index 0000000..2d8ddd9 --- /dev/null +++ b/examples/custom_personality.py @@ -0,0 +1,58 @@ +""" +Example: Creating a custom 'roast' personality for EmotiGrad. + +Shows how to write a Personality callable and pass it to EmotionalOptimizer. +""" + +import torch + +from emotigrad import EmotionalOptimizer + +# --- Custom Personality ------------------------------------------------------- + + +def roast(loss, prev_loss, step): + """A sarcastic personality that roasts the model's progress.""" + if prev_loss is None: + return f"πŸ”₯ Step {step}: New model? Cute. Let's watch it struggle. (avg loss {loss:.4f})" + + if loss < prev_loss: + return f"😏 Step {step}: Look at you improving! Honestly shocked. ({prev_loss:.4f} β†’ {loss:.4f})" + + if loss > prev_loss: + return ( + f"πŸ™ƒ Step {step}: Nice, you made it worse. " + f"({prev_loss:.4f} β†’ {loss:.4f}). Truly groundbreaking." + ) + + return f"🀨 Step {step}: No change. Riveting." + + +# --- Training Loop ------------------------------------------------------------ + + +def main(): + model = torch.nn.Linear(5, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + # Use the custom roast personality + opt = EmotionalOptimizer( + base_opt, + personality=roast, # <-- pass the callable directly + message_every=3, # roast based on averaged loss every 3 steps + ) + + for step in range(12): + x = torch.randn(32, 5) + y = torch.randn(32, 1) + + preds = model(x) + loss = (preds - y).pow(2).mean() + + opt.zero_grad() + loss.backward() + opt.step(loss=loss.item()) + + +if __name__ == "__main__": + main() diff --git a/examples/mnist_training.py b/examples/mnist_training.py new file mode 100644 index 0000000..0d3204f --- /dev/null +++ b/examples/mnist_training.py @@ -0,0 +1,48 @@ +import torch +from torch import nn +from torch.utils.data import DataLoader +from torchvision import datasets, transforms + +from emotigrad import EmotionalOptimizer + +# --- Data --- +transform = transforms.Compose( + [ + transforms.ToTensor(), + ] +) + +train_dataset = datasets.MNIST( + root="data", + train=True, + transform=transform, + download=True, +) +train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) + +# --- Model --- +model = nn.Sequential( + nn.Flatten(), + nn.Linear(28 * 28, 128), + nn.ReLU(), + nn.Linear(128, 10), +) + +# --- Optimizer --- +base_opt = torch.optim.Adam(model.parameters(), lr=1e-3) +opt = EmotionalOptimizer(base_opt, personality="wholesome", message_every=100) + +criterion = nn.CrossEntropyLoss() + +# --- Training Loop --- +for epoch in range(1, 3): + for step, (x, y) in enumerate(train_loader): + preds = model(x) + loss = criterion(preds, y) + + opt.zero_grad() + loss.backward() + opt.step(loss=loss.item()) # <-- triggers emotional output! + + if step % 200 == 0: + print(f"[Epoch {epoch}] step={step}, loss={loss.item():.4f}") diff --git a/pyproject.toml b/pyproject.toml index 84fee8b..cf20ab7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,8 +6,19 @@ authors = [{ name = "Jordan Sinclair" }] readme = "README.md" license = { text = "MIT" } requires-python = ">=3.9" + +keywords = ["pytorch", "optimizer", "deep-learning", "training", "fun"] +classifiers = [ + "Programming Language :: Python :: 3", + "License :: OSI Approved :: MIT License", + "Intended Audience :: Developers", + "Intended Audience :: Science/Research", + "Topic :: Scientific/Engineering :: Artificial Intelligence", +] + dependencies = [ - "torch" + "torch", + "numpy", ] [project.optional-dependencies] diff --git a/src/emotigrad/base.py b/src/emotigrad/base.py index 50ee84f..e0b2a7a 100644 --- a/src/emotigrad/base.py +++ b/src/emotigrad/base.py @@ -1,14 +1,107 @@ -# Imports for the EmotionalOptimizer class -import torch +from __future__ import annotations +from dataclasses import dataclass +from typing import Optional, Union +from torch.optim import Optimizer + +from .types import Personality + +PersonalityLike = Union[str, Personality] + + +@dataclass class EmotionalOptimizer: - def __init__( - self, optimizer: torch.optim.Optimizer, personality: str = "wholesome" - ): - self.optimizer = optimizer - self.personality = personality - - def step(self, *args, **kwargs): - # TODO: add emotional logging later - return self.optimizer.step(*args, **kwargs) + """Wrap a PyTorch optimizer and add emotional feedback. + + This is intentionally lightweight for the first version: + - forwards all calls to the underlying optimizer + - tracks step count and previous loss + - optionally emits text messages via a personality callable + """ + + optimizer: Optimizer + personality: PersonalityLike = "wholesome" + enabled: bool = True + print_fn: callable = print # allows tests / users to override output + message_every: int = 1 # Number of steps between messages + + def __post_init__(self) -> None: + self._step: int = 0 + self._prev_loss: Optional[float] = None + + self._step: int = 0 + self._prev_avg_loss: Optional[float] = None + + # For averaging over the last N steps + self._block_loss_sum: float = 0.0 + self._block_loss_count: int = 0 + + # Resolve personality if given as a string + if isinstance(self.personality, str): + # Lazy import to avoid circular imports + from .personalities import get_personality + + self.personality = get_personality(self.personality) + + def step(self, loss: Optional[float] = None, *args, **kwargs): + """Perform an optimization step and optionally emit emotional feedback. + + Parameters + ---------- + loss: + Current scalar loss value (e.g., loss.item()). + If None, no emotional message is generated for this step. + *args, **kwargs: + Forwarded to the underlying optimizer's step() method. + """ + result = self.optimizer.step(*args, **kwargs) + self._step += 1 + + # Track losses for averaging + if loss is not None: + self._block_loss_sum += float(loss) + self._block_loss_count += 1 + + # Decide whether to emit feedback + if ( + self.enabled + and loss is not None + and self.message_every > 0 + and (self._step % self.message_every == 0) + and self._block_loss_count > 0 + ): + current_avg = self._block_loss_sum / self._block_loss_count + + try: + # message = self.personality(loss, self._prev_loss, self._step) + message = self.personality( + current_avg, + self._prev_avg_loss, + self._step, + ) + except Exception: + # Personality logic should never break training. + message = None + + if message: + self.print_fn(message) + + # Prepare for next block + self._prev_avg_loss = current_avg + self._block_loss_sum = 0.0 + self._block_loss_count = 0 + + return result + + def zero_grad(self, *args, **kwargs): + """Forward zero_grad to the underlying optimizer.""" + return self.optimizer.zero_grad(*args, **kwargs) + + @property + def step_count(self) -> int: + return self._step + + @property + def previous_loss(self) -> Optional[float]: + return self._prev_loss diff --git a/src/emotigrad/personalities.py b/src/emotigrad/personalities.py new file mode 100644 index 0000000..21eaf67 --- /dev/null +++ b/src/emotigrad/personalities.py @@ -0,0 +1,115 @@ +# src/emotigrad/personalities.py +from __future__ import annotations + +from typing import Dict, List, Optional + +from .types import Personality + + +# --- Built-in personalities ------------------------------------------------- +def _default_personality( + loss: float, + prev_loss: Optional[float], + step: int, +) -> Optional[str]: + """Very minimal 'wholesome' personality for early versions. + + You can expand this or move it into a dedicated personalities module later. + """ + if prev_loss is None: + return f"✨ Starting our journey! Initial loss: {loss:.4f}" + + if loss < prev_loss: + return f"πŸ’– Nice! Loss improved from {prev_loss:.4f} to {loss:.4f}." + + if loss > prev_loss: + return ( + f"🌱 It's okay! Loss went from {prev_loss:.4f} to {loss:.4f}. " + "Learning isn't always linear." + ) + + return None + + +def wholesome(loss: float, prev_loss: Optional[float], step: int) -> Optional[str]: + if prev_loss is None: + return f"✨ Let's get started! Initial loss: {loss:.4f}" + + if loss < prev_loss: + return f"πŸ’– Nice! Loss improved from {prev_loss:.4f} to {loss:.4f}." + + if loss > prev_loss: + return ( + f"🌱 It's okay! Loss went from {prev_loss:.4f} to {loss:.4f}. " + "Learning isn't always monotonic." + ) + + return None # no message if unchanged + + +def sassy(loss: float, prev_loss: Optional[float], step: int) -> Optional[str]: + if prev_loss is None: + return "πŸ˜’ Fine, let's see what you've got." + + if loss > prev_loss: + return f"πŸ™„ Bold move: loss got worse ({prev_loss:.4f} β†’ {loss:.4f})." + + if loss < prev_loss: + return f"πŸ‘ About time: {prev_loss:.4f} β†’ {loss:.4f}." + + return "🀨 Exactly the same? Interesting choice." + + +class QuietPersonality: + """Example of a personality implemented as a class with state.""" + + def __init__(self, every_n_steps: int = 10) -> None: + self.every_n_steps = every_n_steps + + def __call__( + self, loss: float, prev_loss: Optional[float], step: int + ) -> Optional[str]: + if step % self.every_n_steps != 0: + return None + return f"πŸ”Ž Step {step}: current loss {loss:.4f}" + + +# --- Registry --------------------------------------------------------------- + + +_PERSONALITY_REGISTRY: Dict[str, Personality] = { + "wholesome": wholesome, + "sassy": sassy, + "quiet": QuietPersonality(), # instance is fine; it's still callable +} + + +def register_personality( + name: str, + personality: Personality, + *, + overwrite: bool = False, +) -> None: + """Register a new personality under a string name. + + Users can call this to add their own personalities. + """ + key = name.lower() + if not overwrite and key in _PERSONALITY_REGISTRY: + raise ValueError(f"Personality '{name}' is already registered.") + _PERSONALITY_REGISTRY[key] = personality + + +def get_personality(name: str) -> Personality: + """Look up a personality by name, raising KeyError if not found.""" + key = name.lower() + try: + return _PERSONALITY_REGISTRY[key] + except KeyError as exc: + available = ", ".join(sorted(_PERSONALITY_REGISTRY.keys())) + raise KeyError(f"Unknown personality '{name}'. Available: {available}") from exc + + +def list_personalities() -> List[str]: + """Return a sorted list of available personality names.""" + return sorted(_PERSONALITY_REGISTRY.keys()) diff --git a/src/emotigrad/types.py b/src/emotigrad/types.py new file mode 100644 index 0000000..789270a --- /dev/null +++ b/src/emotigrad/types.py @@ -0,0 +1,27 @@ +from typing import Optional, Protocol, runtime_checkable + + +@runtime_checkable +class Personality(Protocol): + """Protocol for personality callables.""" + + def __call__( + self, + loss: float, + prev_loss: Optional[float], + step: int, + ) -> Optional[str]: + """Generate an emotional message based on the optimization state. + + Parameters + ---------- + loss: + Current scalar loss value. + prev_loss: + Previous scalar loss value, or None if this is the first step. + step: + Current optimization step count (starting from 1). + Returns: + An optional string message to emit. + """ + ... diff --git a/tests/test_emotional_optimizer_basic.py b/tests/test_emotional_optimizer_basic.py new file mode 100644 index 0000000..ff7f4e2 --- /dev/null +++ b/tests/test_emotional_optimizer_basic.py @@ -0,0 +1,157 @@ +import torch + +from emotigrad import EmotionalOptimizer + + +class DummyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.linear = torch.nn.Linear(2, 1) + + def forward(self, x): + return self.linear(x) + + +def test_emotional_optimizer_wraps_optimizer_and_steps(): + model = DummyModel() + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + emo_opt = EmotionalOptimizer(base_opt, enabled=False) + + x = torch.randn(4, 2) + y = torch.randn(4, 1) + + # compute a loss + preds = model(x) + loss = (preds - y).pow(2).mean() + + # standard training step using EmotionalOptimizer + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + # parameters should have changed + with torch.no_grad(): + preds_after = model(x) + new_loss = (preds_after - y).pow(2).mean() + + assert emo_opt.step_count == 1 + # We don't assert that loss strictly decreases, but + # we at least verify that something happened. + assert new_loss != loss + + +def test_emotional_optimizer_calls_personality_when_enabled(): + model = DummyModel() + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + messages = [] + + def fake_print(msg: str): + messages.append(msg) + + def fake_personality(loss, prev_loss, step): + return f"step={step}, loss={loss}, prev={prev_loss}" + + emo_opt = EmotionalOptimizer( + base_opt, + personality=fake_personality, + enabled=True, + print_fn=fake_print, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + assert len(messages) == 1 + assert "step=1" in messages[0] + + +# Confirm personality is NOT called when enabled=False +def test_emotional_optimizer_personality_not_called_when_disabled(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + called = [] + + def fake_personality(loss, prev_loss, step): + called.append((loss, prev_loss, step)) + return None + + emo_opt = EmotionalOptimizer( + base_opt, + personality=fake_personality, + enabled=False, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + assert len(called) == 0 + + +# Confirm personality is called when loss is provided and enabled=True +def test_smoke_emotional_optimizer_personality_called(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + called = [] + + def fake_personality(loss, prev_loss, step): + called.append((loss, prev_loss, step)) + return None + + emo_opt = EmotionalOptimizer( + base_opt, + personality=fake_personality, + enabled=True, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + assert len(called) == 1 + + +# Confirm exceptions inside personalities do not crash step() and are safely swallowed +def test_emotional_optimizer_personality_exceptions_handled(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + def faulty_personality(loss, prev_loss, step): + raise RuntimeError("Deliberate failure inside personality") + + emo_opt = EmotionalOptimizer( + base_opt, + personality=faulty_personality, + enabled=True, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + + # This should not raise, despite the personality raising internally + emo_opt.step(loss=loss.item()) diff --git a/tests/test_message_every.py b/tests/test_message_every.py new file mode 100644 index 0000000..6f6b6b9 --- /dev/null +++ b/tests/test_message_every.py @@ -0,0 +1,45 @@ +import torch + +from emotigrad import EmotionalOptimizer + + +def test_message_every_uses_block_average_loss(): + model = torch.nn.Linear(1, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + messages = [] + + def fake_print(msg: str): + messages.append(msg) + + # This personality simply prints the loss it's given + def echo_personality(loss, prev_loss, step): + return f"loss={loss:.2f}, prev={None if prev_loss is None else round(prev_loss, 2)}" + + opt = EmotionalOptimizer( + base_opt, + personality=echo_personality, + print_fn=fake_print, + message_every=3, + ) + + # Provide known loss values + # Block 1 = [1.0, 2.0, 3.0] β†’ average = 2.0 + # Block 2 = [4.0, 6.0, 8.0] β†’ average = 6.0 + + losses = [1.0, 2.0, 3.0, 4.0, 6.0, 8.0] + + # Dummy step (no gradients needed) + for loss in losses: + opt.optimizer.zero_grad() + opt.step(loss=loss) + + assert len(messages) == 2, "Should emit once per block of 3 steps." + + # First message + assert "loss=2.00" in messages[0] + assert "prev=None" in messages[0] + + # Second message + assert "loss=6.00" in messages[1] + assert "prev=2.0" in messages[1] diff --git a/tests/test_smoke.py b/tests/test_smoke.py index 30dc425..d8bf62c 100644 --- a/tests/test_smoke.py +++ b/tests/test_smoke.py @@ -15,3 +15,96 @@ def test_smoke_emotional_optimizer_wraps_optimizer(): loss.backward() emo_opt.step() + + +# Verify that personalities are called correctly and messages are routed through the configured print_fn. +def test_smoke_emotional_optimizer_personality(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + messages = [] + + def fake_print(msg: str): + messages.append(msg) + + emo_opt = EmotionalOptimizer( + base_opt, + personality="sassy", + enabled=True, + print_fn=fake_print, + ) + + x = torch.randn(4, 2) + y = torch.randn(4, 1) + + # First step + loss1 = (model(x) - y).pow(2).mean() + loss1.backward() + emo_opt.step(loss=loss1.item()) + + # Second step with worse loss + with torch.no_grad(): + model.weight += 0.5 # make loss worse + loss2 = (model(x) - y).pow(2).mean() + loss2.backward() + emo_opt.step(loss=loss2.item()) + + assert len(messages) == 2 + + +# Confirm personality is called when loss is provided and enabled=True +def test_smoke_emotional_optimizer_personality_called(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + called = [] + + def fake_personality(loss, prev_loss, step): + called.append((loss, prev_loss, step)) + return None + + emo_opt = EmotionalOptimizer( + base_opt, + personality=fake_personality, + enabled=True, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + assert len(called) == 1 + + +# Confirm personality is NOT called when enabled=False +def test_smoke_emotional_optimizer_personality_not_called_when_disabled(): + model = torch.nn.Linear(2, 1) + base_opt = torch.optim.SGD(model.parameters(), lr=0.1) + + called = [] + + def fake_personality(loss, prev_loss, step): + called.append((loss, prev_loss, step)) + return None + + emo_opt = EmotionalOptimizer( + base_opt, + personality=fake_personality, + enabled=False, + ) + + x = torch.randn(2, 2) + y = torch.randn(2, 1) + preds = model(x) + loss = (preds - y).pow(2).mean() + + emo_opt.zero_grad() + loss.backward() + emo_opt.step(loss=loss.item()) + + assert len(called) == 0