Skip to content

StrataNet Quick Start

StrataNet is a PyTorch neural network architecture. It requires torch>=2.0.

Train from OHLCV data

from strata import StrataNet, StrataNetConfig, StrataNetTrainer, StrataNetDataset

# Your candles: list of OHLCV dicts, oldest → newest
candles = [
    {"open": 150.0, "high": 151.2, "low": 149.5, "close": 150.8, "volume": 1_200_000},
    # ... minimum seq_len+1 candles
]

# Labels auto-generated by STRATA state machine (no manual labeling)
dataset = StrataNetDataset.from_candles(candles, seq_len=30, asset="AAPL")
trainer = StrataNetTrainer(asset="AAPL")
model   = trainer.train(dataset, epochs=30)

model.save("aapl_net.pt")   # standard PyTorch format

Load and predict

import torch
from strata import StrataNet
from strata.net_trainer import _normalise_window

model   = StrataNet.load("aapl_net.pt")
window  = candles[-30:]   # last 30 candles
x       = torch.tensor(_normalise_window(window), dtype=torch.float32)
result  = model.predict_action(x)

print(result)
# {
#   "action":     "LONG",
#   "confidence": 0.71,
#   "regime":     "TRENDING",
#   "state": {
#       "bias":        0.82,    # [-1, 1]  directional conviction
#       "momentum":    0.45,    # [0, 1]   breakout energy
#       "trap_risk":   0.18,    # [0, 1]   adverse selection risk
#       "uncertainty": 0.31,    # [0, 1]   volatility ambiguity
#   }
# }

Load pretrained from Hugging Face

from huggingface_hub import hf_hub_download
from strata import StrataNet

path  = hf_hub_download(repo_id="emylton/strata-net", filename="aapl_net.pt")
model = StrataNet.load(path)

Available pretrained: aapl_net.pt, tsla_net.pt, spy_net.pt, nvda_net.pt, qqq_net.pt, btc_net.pt

Fine-tune pretrained

from strata import StrataNetTrainer, StrataNetDataset

dataset = StrataNetDataset.from_candles(my_candles, seq_len=30, asset="AAPL")
trainer = StrataNetTrainer(model=pretrained_model, lr=1e-4)  # lower LR
model   = trainer.train(dataset, epochs=10)

Custom config

cfg = StrataNetConfig(
    asset       = "TSLA",
    seq_len     = 60,     # longer lookback
    embed_dim   = 64,     # larger embedding
    core_expand = 32,
    head_dim    = 32,
    dropout     = 0.15,
)
model = StrataNet(cfg)
print(model.summary())