Download daisychain/example_task.py from DaisyChainAI/DaisyChain-Train: direct link, hf CLI and curl.
- Browser
- Download file 920 Bytes
-
https://huggingface.co/DaisyChainAI/DaisyChain-Train/resolve/2841ec9218506e90ea508ac6b71bb8c740983db5/daisychain/example_task.py
- Command line
-
hf download hf://DaisyChainAI/DaisyChain-Train@2841ec9218506e90ea508ac6b71bb8c740983db5/daisychain/example_task.py
-
curl -L -o example_task.py https://huggingface.co/DaisyChainAI/DaisyChain-Train/resolve/2841ec9218506e90ea508ac6b71bb8c740983db5/daisychain/example_task.py
920 Bytes
| """The default example task: fit a small MLP to a synthetic function. | |
| It exists so `daisychain-train` runs out of the box and you can confirm the | |
| cluster works end to end. Replace it with your own task (see docs/CUSTOM_TASK.md) | |
| -- copy this file, change build_model / sample / loss, and set DAISY_TASK. | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| class ExampleTask: | |
| def __init__(self): | |
| # fixed target so every node's shard is consistent | |
| g = torch.Generator().manual_seed(1234) | |
| self.W = torch.randn(8, 1, generator=g) | |
| def build_model(self): | |
| torch.manual_seed(0) # identical init on every node | |
| return nn.Sequential(nn.Linear(8, 32), nn.ReLU(), nn.Linear(32, 1)) | |
| def sample(self, n): | |
| X = torch.randn(n, 8) | |
| return X, X @ self.W + 0.05 * torch.randn(n, 1) | |
| def loss(self, model, X, y): | |
| return nn.functional.mse_loss(model(X), y) | |