Skip to content

Commit a175f75

Browse files
committed
Changed location for the hedge_net_tx.pth, fixed pre-commit issues in evaluate.py
1 parent 2ef3b2f commit a175f75

2 files changed

Lines changed: 46 additions & 35 deletions

File tree

ml/evaluate.py

Lines changed: 43 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -1,84 +1,93 @@
1-
# ml/evaluate.py
1+
"""
2+
Evaluating HedgeNet.
3+
4+
This module contains implementation of HedgeNet evaluation pipeline.
5+
"""
6+
27
import torch
3-
import numpy as np
8+
49
from ml.config import GBMConfig, HedgingConfig, StressTestConfig
5-
from ml.sim.gbm import simulate_gbm
10+
from ml.metrics.pnl import compute_pnl_with_tx
611
from ml.models.hedge_net import HedgeNet
7-
from ml.metrics.pnl import compute_pnl_with_tx, decompose_pnl
12+
from ml.sim.gbm import simulate_gbm
813

9-
def load_model(hidden_dim, path="hedge_net_tx.pth", device="cpu"):
14+
15+
def load_model(hidden_dim, path="artifacts/hedge_net_tx.pth", device="cpu"):
16+
"""Load HedgeNet model from the file."""
1017
net = HedgeNet(hidden_dim).to(device)
1118
net.load_state_dict(torch.load(path, weights_only=True))
1219
net.eval()
1320
return net
1421

15-
def prepare_inputs_for_model(S, K, T, M, device='cpu'):
16-
t_grid = torch.linspace(0, T - T/M, M, device=device)
22+
23+
def prepare_inputs_for_model(S, K, T, M, device="cpu"):
24+
"""Prepare inputs for HedgeNet evaluation."""
25+
t_grid = torch.linspace(0, T - T / M, M, device=device)
1726
tau = T - t_grid
1827
N = S.size(0)
1928
tau_batch = tau.unsqueeze(0).expand(N, -1).reshape(-1)
2029
moneyness_batch = (S[:, :-1] / K).reshape(-1)
2130
return tau_batch, moneyness_batch, N, M
2231

32+
2333
def run_stress_test():
34+
"""Run some testing scenarios on hedging with HedgeNet."""
2435
base_gbm = GBMConfig()
2536
base_hedge = HedgingConfig()
2637
stress_cfg = StressTestConfig()
2738

2839
net = load_model(base_hedge.hidden_dim, device=base_hedge.device)
2940

30-
results = {
31-
'sigma': [],
32-
'lambda_tx': [],
33-
'M': [],
34-
'mean_abs_pnl': [],
35-
'std_pnl': []
36-
}
41+
results = {"sigma": [], "lambda_tx": [], "M": [], "mean_abs_pnl": [], "std_pnl": []}
3742

3843
# Vary sigma
3944
for sigma in stress_cfg.sigma_vals:
4045
gbm = GBMConfig(S0=base_gbm.S0, sigma=sigma, T=base_gbm.T, N=5000, M=base_gbm.M)
4146
S = simulate_gbm(**gbm.__dict__, device=base_hedge.device).float()
42-
tau_flat, moneyness_flat, N, M = prepare_inputs_for_model(S, base_hedge.K, gbm.T, gbm.M)
47+
tau_flat, moneyness_flat, N, M = prepare_inputs_for_model(
48+
S, base_hedge.K, gbm.T, gbm.M
49+
)
4350
with torch.no_grad():
4451
phi_flat = net(tau_flat, moneyness_flat)
4552
phi = phi_flat.reshape(N, M)
4653
pnl = compute_pnl_with_tx(S, base_hedge.K, phi, base_hedge.lambda_tx)
47-
results['sigma'].append(sigma)
48-
results['lambda_tx'].append(base_hedge.lambda_tx)
49-
results['M'].append(gbm.M)
50-
results['mean_abs_pnl'].append(pnl.abs().mean().item())
51-
results['std_pnl'].append(pnl.std().item())
54+
results["sigma"].append(sigma)
55+
results["lambda_tx"].append(base_hedge.lambda_tx)
56+
results["M"].append(gbm.M)
57+
results["mean_abs_pnl"].append(pnl.abs().mean().item())
58+
results["std_pnl"].append(pnl.std().item())
5259

5360
# Vary lambda_tx
5461
for lam in stress_cfg.lambda_vals:
5562
S = simulate_gbm(**base_gbm.__dict__, device=base_hedge.device).float()
56-
tau_flat, moneyness_flat, N, M = prepare_inputs_for_model(S, base_hedge.K, base_gbm.T, base_gbm.M)
63+
tau_flat, moneyness_flat, N, M = prepare_inputs_for_model(
64+
S, base_hedge.K, base_gbm.T, base_gbm.M
65+
)
5766
with torch.no_grad():
5867
phi_flat = net(tau_flat, moneyness_flat)
5968
phi = phi_flat.reshape(N, M)
6069
pnl = compute_pnl_with_tx(S, base_hedge.K, phi, lam)
61-
results['sigma'].append(base_gbm.sigma)
62-
results['lambda_tx'].append(lam)
63-
results['M'].append(base_gbm.M)
64-
results['mean_abs_pnl'].append(pnl.abs().mean().item())
65-
results['std_pnl'].append(pnl.std().item())
70+
results["sigma"].append(base_gbm.sigma)
71+
results["lambda_tx"].append(lam)
72+
results["M"].append(base_gbm.M)
73+
results["mean_abs_pnl"].append(pnl.abs().mean().item())
74+
results["std_pnl"].append(pnl.std().item())
6675

6776
# Vary M (rebalancing frequency)
6877
for M in stress_cfg.M_vals:
6978
gbm = GBMConfig(S0=base_gbm.S0, sigma=base_gbm.sigma, T=base_gbm.T, N=5000, M=M)
7079
S = simulate_gbm(**gbm.__dict__, device=base_hedge.device).float()
71-
tau_flat, moneyness_flat, N, M_actual = prepare_inputs_for_model(S, base_hedge.K, gbm.T, gbm.M)
80+
tau_flat, moneyness_flat, N, M_actual = prepare_inputs_for_model(
81+
S, base_hedge.K, gbm.T, gbm.M
82+
)
7283
with torch.no_grad():
7384
phi_flat = net(tau_flat, moneyness_flat)
7485
phi = phi_flat.reshape(N, M_actual)
7586
pnl = compute_pnl_with_tx(S, base_hedge.K, phi, base_hedge.lambda_tx)
76-
results['sigma'].append(base_gbm.sigma)
77-
results['lambda_tx'].append(base_hedge.lambda_tx)
78-
results['M'].append(M)
79-
results['mean_abs_pnl'].append(pnl.abs().mean().item())
80-
results['std_pnl'].append(pnl.std().item())
87+
results["sigma"].append(base_gbm.sigma)
88+
results["lambda_tx"].append(base_hedge.lambda_tx)
89+
results["M"].append(M)
90+
results["mean_abs_pnl"].append(pnl.abs().mean().item())
91+
results["std_pnl"].append(pnl.std().item())
8192

8293
return results
83-
84-

ml/train.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
python -m ml.train --data_type_model=heston
1717
"""
1818
import argparse
19+
import os
1920

2021
import torch
2122

@@ -79,7 +80,8 @@ def main(args):
7980
if epoch % 100 == 0:
8081
print(f"Epoch {epoch}, Loss: {loss.item():.4f}")
8182

82-
torch.save(net.state_dict(), "hedge_net_tx.pth")
83+
os.makedirs("artifacts", exist_ok=True)
84+
torch.save(net.state_dict(), "artifacts/hedge_net_tx.pth")
8385
print("Model saved!")
8486

8587

0 commit comments

Comments
 (0)