File size: 17,494 Bytes
fa2c1e2
9fa5bf3
fa2c1e2
 
9fa5bf3
 
fa2c1e2
 
6c3178c
9fa5bf3
 
 
 
fa2c1e2
 
 
 
 
 
 
 
 
 
 
0122c63
fa2c1e2
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
9fa5bf3
 
6c3178c
9fa5bf3
 
 
fa2c1e2
9fa5bf3
 
fa2c1e2
 
 
 
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
 
 
 
 
91ba9ec
fa2c1e2
 
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
9fa5bf3
fa2c1e2
 
 
 
 
9fa5bf3
fa2c1e2
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9fa5bf3
 
 
 
 
fa2c1e2
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
 
95c00cd
 
 
 
 
 
 
 
fa2c1e2
95c00cd
fa2c1e2
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
91ba9ec
 
fa2c1e2
 
9fa5bf3
fa2c1e2
 
492e6ca
9fa5bf3
 
 
6c3178c
9fa5bf3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fa2c1e2
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
9fa5bf3
6c3178c
 
 
 
 
 
 
fa2c1e2
 
9fa5bf3
fa2c1e2
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
 
 
 
6c3178c
fa2c1e2
 
 
 
 
 
 
 
91ba9ec
b294be0
 
 
 
 
 
 
 
9fa5bf3
b294be0
 
 
 
 
9fa5bf3
 
 
 
 
 
fa2c1e2
 
 
 
 
 
 
9fa5bf3
 
492e6ca
 
 
 
 
9fa5bf3
 
 
 
 
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
9fa5bf3
fa2c1e2
 
 
9fa5bf3
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
 
9fa5bf3
fa2c1e2
9fa5bf3
fa2c1e2
 
 
 
9fa5bf3
fa2c1e2
 
 
 
 
 
 
 
 
 
 
9fa5bf3
fa2c1e2
9fa5bf3
6c3178c
9fa5bf3
fa2c1e2
9fa5bf3
 
 
 
 
 
fa2c1e2
 
 
9fa5bf3
 
 
fa2c1e2
 
492e6ca
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
"""
Multi-Symbol Signal Dashboard (Educational) — Hugging Face Space
Signal-only display. No auto-execution. User manually places trades on MT5 demo.

XAUUSD: PPO RL model (JonusNattapong/Reinforcement-Learning-for-Gold-Trading-Model),
trained strictly on 15-minute XAUUSD bars — inference here uses native 15min data
to match training distribution (do not feed other timeframes into this model).

BTCUSD: no validated BTC-specific RL model was found, so BTC uses a transparent
rule-based technical signal (RSI + EMA trend + momentum) instead of misapplying
the gold-only PPO model to a different asset class.

Data: TwelveData API (15min interval, free tier: 800 credits/day, 8/min)
"""

import os
import pickle
from datetime import datetime, timezone

import numpy as np
import pandas as pd
import plotly.graph_objects as go
import requests
import gradio as gr
import spaces
from huggingface_hub import hf_hub_download
from stable_baselines3 import PPO

# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------

TWELVEDATA_API_KEY = os.environ.get("TWELVEDATA_API_KEY", "")
INTERVAL = "15min"  # must match PPO model training cadence — do not change
OUTPUT_SIZE = 200  # enough bars for rolling windows (32-period vol, etc.)

SYMBOLS = {
    "XAUUSD": {"td_symbol": "XAU/USD", "engine": "ppo", "decimals": 2},
    "BTCUSD": {"td_symbol": "BTC/USD", "engine": "rules", "decimals": 2},
}
DEFAULT_SYMBOL = "XAUUSD"

MODEL_REPO = "JonusNattapong/Reinforcement-Learning-for-Gold-Trading-Model"
MODEL_FILE = "ppo_xauusd.zip"
VECNORM_FILE = "vecnormalize.pkl"

FEATURE_COLS = [
    "log_return", "hl_range", "body", "atr14",
    "rsi14", "ema_diff", "volatility", "tod_sin", "tod_cos",
]

TP_MULTIPLES = [1.5, 3.0, 5.3]
ACTION_LABELS = {0: "HOLD", 1: "BUY", 2: "SELL"}
ACTION_COLORS = {"BUY": "#16c784", "SELL": "#ea3943", "HOLD": "#f3ba2f"}

# ---------------------------------------------------------------------------
# PPO model loading (cached — used only for XAUUSD)
# ---------------------------------------------------------------------------

_model = None
_obs_mean = None
_obs_var = None
_clip_obs = 10.0
_epsilon = 1e-8


def load_model():
    global _model, _obs_mean, _obs_var, _clip_obs
    if _model is not None:
        return
    model_path = hf_hub_download(repo_id=MODEL_REPO, filename=MODEL_FILE)
    _model = PPO.load(model_path, device="cpu")
    vecnorm_path = hf_hub_download(repo_id=MODEL_REPO, filename=VECNORM_FILE)
    with open(vecnorm_path, "rb") as f:
        vecnorm = pickle.load(f)
    _obs_mean = np.asarray(vecnorm.obs_rms.mean, dtype=np.float32)
    _obs_var = np.asarray(vecnorm.obs_rms.var, dtype=np.float32)
    _clip_obs = float(getattr(vecnorm, "clip_obs", 10.0))


# ---------------------------------------------------------------------------
# Data fetching
# ---------------------------------------------------------------------------

def fetch_ohlc(symbol_key: str) -> pd.DataFrame:
    if not TWELVEDATA_API_KEY:
        raise RuntimeError(
            "TWELVEDATA_API_KEY not set. Add it as a Space secret "
            "(Settings → Repository secrets)."
        )
    td_symbol = SYMBOLS[symbol_key]["td_symbol"]
    url = "https://api.twelvedata.com/time_series"
    params = {
        "symbol": td_symbol,
        "interval": INTERVAL,
        "outputsize": OUTPUT_SIZE,
        "timezone": "UTC",
        "order": "asc",
        "apikey": TWELVEDATA_API_KEY,
    }
    resp = requests.get(url, params=params, timeout=15)
    data = resp.json()
    if "values" not in data:
        raise RuntimeError(f"TwelveData error: {data.get('message', data)}")

    df = pd.DataFrame(data["values"])
    df["datetime"] = pd.to_datetime(df["datetime"], utc=True)
    for col in ["open", "high", "low", "close"]:
        df[col] = pd.to_numeric(df[col])
    df["volume"] = pd.to_numeric(df.get("volume", 0))
    df = df.sort_values("datetime").set_index("datetime")
    return df


# ---------------------------------------------------------------------------
# Feature engineering — mirrors rl_gold_trading/features.py exactly
# ---------------------------------------------------------------------------

def _rsi(series: pd.Series, period: int = 14) -> pd.Series:
    delta = series.diff()
    gain = delta.clip(lower=0.0)
    loss = -delta.clip(upper=0.0)
    avg_gain = gain.rolling(period).mean()
    avg_loss = loss.rolling(period).mean()
    rs = avg_gain / (avg_loss.replace(0.0, np.nan))
    rsi = 100.0 - (100.0 / (1.0 + rs))
    return rsi.fillna(50.0)


def add_features(df: pd.DataFrame) -> pd.DataFrame:
    df = df.copy()
    df["log_return"] = np.log(df["close"]).diff()
    df["hl_range"] = (df["high"] - df["low"]) / df["close"]
    df["body"] = (df["close"] - df["open"]) / df["close"]
    prev_close = df["close"].shift(1)
    tr = pd.concat(
        [
            (df["high"] - df["low"]),
            (df["high"] - prev_close).abs(),
            (df["low"] - prev_close).abs(),
        ],
        axis=1,
    ).max(axis=1)
    df["atr14"] = tr.rolling(14).mean() / df["close"]
    df["rsi14"] = _rsi(df["close"], 14) / 100.0
    ema12 = df["close"].ewm(span=12, adjust=False).mean()
    ema26 = df["close"].ewm(span=26, adjust=False).mean()
    macd_line = ema12 - ema26
    macd_signal = macd_line.ewm(span=9, adjust=False).mean()
    df["macd_hist"] = macd_line - macd_signal
    ema20 = df["close"].ewm(span=20, adjust=False).mean()
    ema50 = df["close"].ewm(span=50, adjust=False).mean()
    df["ema_diff"] = (ema20 - ema50) / df["close"]
    df["volatility"] = df["log_return"].rolling(32).std()
    minutes = df.index.hour * 60 + df.index.minute
    day_frac = minutes / (24 * 60)
    df["tod_sin"] = np.sin(2 * np.pi * day_frac)
    df["tod_cos"] = np.cos(2 * np.pi * day_frac)
    df = df.dropna(subset=FEATURE_COLS + ["macd_hist"])
    return df


# ---------------------------------------------------------------------------
# PPO inference (XAUUSD only)
# ---------------------------------------------------------------------------

def build_observation(feat_row: pd.Series) -> np.ndarray:
    """
    Uses the training-set mean for the 4 state dims (position, unrealized PnL,
    daily PnL, trades today) instead of raw zeros. Raw [0,0,0,0] represents an
    edge case rarely seen in training (only at the very first step of an
    episode) and tends to saturate the policy toward HOLD. Feeding the mean
    normalizes to ~0 (neutral) and lets the decision hinge on market features
    instead, which better matches how the model behaves mid-session.
    """
    feat = feat_row[FEATURE_COLS].to_numpy(dtype=np.float32)
    state = np.asarray(_obs_mean[9:13], dtype=np.float32) if _obs_mean is not None else np.zeros(4, dtype=np.float32)
    return np.concatenate([feat, state])


def normalize_obs(obs: np.ndarray) -> np.ndarray:
    norm = (obs - _obs_mean) / np.sqrt(_obs_var + _epsilon)
    return np.clip(norm, -_clip_obs, _clip_obs).astype(np.float32)


def predict_ppo(feat_row: pd.Series):
    obs = build_observation(feat_row)
    norm_obs = normalize_obs(obs)
    import torch
    _model.policy.to("cpu")
    obs_tensor = torch.as_tensor(norm_obs).float().unsqueeze(0)
    with torch.no_grad():
        dist = _model.policy.get_distribution(obs_tensor)
        probs = dist.distribution.probs.numpy()[0]
    action = int(np.argmax(probs))
    confidence = float(probs[action]) * 100.0
    return ACTION_LABELS[action], confidence, probs


# ---------------------------------------------------------------------------
# Rule-based signal (BTCUSD — transparent technical logic, not ML)
# ---------------------------------------------------------------------------

def predict_rules(feat_row: pd.Series):
    """
    Simple, transparent technical vote: RSI + EMA trend + MACD histogram.
    Each indicator casts one vote; majority decides action. Confidence is the
    fraction of indicators agreeing (33/67/100%), not a probabilistic estimate.
    """
    rsi = feat_row["rsi14"] * 100.0  # back to 0-100 scale
    ema_diff = feat_row["ema_diff"]
    macd_hist = feat_row["macd_hist"]

    votes = []
    votes.append("BUY" if rsi < 35 else "SELL" if rsi > 65 else "HOLD")
    votes.append("BUY" if ema_diff > 0 else "SELL")
    votes.append("BUY" if macd_hist > 0 else "SELL")

    buy_votes = votes.count("BUY")
    sell_votes = votes.count("SELL")
    if buy_votes >= 2 and buy_votes > sell_votes:
        action = "BUY"
        confidence = (buy_votes / 3) * 100.0
    elif sell_votes >= 2 and sell_votes > buy_votes:
        action = "SELL"
        confidence = (sell_votes / 3) * 100.0
    else:
        action = "HOLD"
        confidence = 33.0
    return action, confidence


# ---------------------------------------------------------------------------
# SL / TP construction (ATR-based, matches the reference card layout)
# ---------------------------------------------------------------------------

def build_levels(entry: float, atr_frac: float, action: str):
    atr_abs = atr_frac * entry
    r = max(atr_abs * 1.5, entry * 0.001)
    if action == "BUY":
        sl = entry - r
        tps = [entry + r * m for m in TP_MULTIPLES]
    elif action == "SELL":
        sl = entry + r
        tps = [entry - r * m for m in TP_MULTIPLES]
    else:
        sl = None
        tps = []
    return sl, tps, r


# ---------------------------------------------------------------------------
# Chart rendering (dark theme, styled like the reference screenshot)
# ---------------------------------------------------------------------------

def render_chart(df: pd.DataFrame, entry: float, sl, tps, action: str, decimals: int) -> go.Figure:
    fig = go.Figure()
    fig.add_trace(go.Candlestick(
        x=df.index, open=df["open"], high=df["high"], low=df["low"], close=df["close"],
        increasing_line_color="#16c784", decreasing_line_color="#ea3943",
        name="price",
    ))

    fmt = f",.{decimals}f"

    # Always show the current price, even on HOLD (no trade levels to draw).
    fig.add_hline(y=entry, line_color="#e8e8e8", line_width=1, line_dash="dot",
                   annotation_text=f"Current {entry:{fmt}}",
                   annotation_position="right",
                   annotation_font_color="#e8e8e8", annotation_font_size=11)

    if action in ("BUY", "SELL") and sl is not None:
        fig.add_hline(y=entry, line_color="#f3ba2f", line_width=1.5,
                       annotation_text=f"{action} {entry:{fmt}}",
                       annotation_position="right",
                       annotation_font_color="#f3ba2f")
        fig.add_hline(y=sl, line_color="#ea3943", line_width=1.5,
                       annotation_text=f"SL {sl:{fmt}} · -1R",
                       annotation_position="right",
                       annotation_font_color="#ea3943")
        y0, y1 = (entry, sl) if action == "BUY" else (sl, entry)
        fig.add_hrect(y0=min(y0, y1), y1=max(y0, y1), fillcolor="#ea3943", opacity=0.12, line_width=0)

        labels = ["TP1", "TP2", "TP3"]
        for label, tp, mult in zip(labels, tps, TP_MULTIPLES):
            fig.add_hline(y=tp, line_color="#16c784", line_width=1.2, line_dash="dash",
                           annotation_text=f"{label} {tp:{fmt}} · {mult}R",
                           annotation_position="right",
                           annotation_font_color="#16c784")
        y0, y1 = (entry, tps[-1]) if action == "BUY" else (tps[-1], entry)
        fig.add_hrect(y0=min(y0, y1), y1=max(y0, y1), fillcolor="#16c784", opacity=0.10, line_width=0)

    fig.update_layout(
        template="plotly_dark",
        paper_bgcolor="#0b0e11", plot_bgcolor="#0b0e11",
        font=dict(color="#e8e8e8"),
        margin=dict(l=10, r=90, t=10, b=10),
        xaxis_rangeslider_visible=False,
        height=460,
        showlegend=False,
        dragmode="pan",  # single-finger drag pans; pinch-to-zoom works natively on touch
    )
    return fig


# ---------------------------------------------------------------------------
# Main callback
# ---------------------------------------------------------------------------

@spaces.GPU(duration=15)
def _gpu_probe():
    """Trivial decorated function so ZeroGPU hardware detects a GPU-capable
    function at startup. We don't actually need GPU compute (model is tiny
    and always runs on CPU) - this just satisfies the platform requirement
    and lets us catch quota-exceeded errors gracefully instead of crashing."""
    return True


def run_signal(symbol_key: str):
    try:
        _gpu_probe()
    except Exception:
        pass  # ZeroGPU quota exhausted - fine, we run on CPU regardless

    if symbol_key not in SYMBOLS:
        symbol_key = DEFAULT_SYMBOL
    cfg = SYMBOLS[symbol_key]
    decimals = cfg["decimals"]

    raw = fetch_ohlc(symbol_key)
    feat_df = add_features(raw)
    if feat_df.empty:
        raise gr.Error("Not enough bars returned to compute features. Try again shortly.")

    last_row = feat_df.iloc[-1]
    entry = float(last_row["close"])

    if cfg["engine"] == "ppo":
        load_model()
        action, confidence, probs = predict_ppo(last_row)
        engine_label = (
            f"PPO model (15m) · Hold {probs[0]*100:.0f}% · "
            f"Buy {probs[1]*100:.0f}% · Sell {probs[2]*100:.0f}%"
        )
    else:
        action, confidence = predict_rules(last_row)
        engine_label = "Rule-based technical signal (15m) — not ML"

    sl, tps, r = build_levels(entry, float(last_row["atr14"]), action)
    rr_text = f"1:{TP_MULTIPLES[0]}" if action in ("BUY", "SELL") else "—"
    ts = datetime.now(timezone.utc).strftime("%-d %b %Y %H:%M UTC")

    badge_color = ACTION_COLORS.get(action, "#8a8f98")
    fmt = f",.{decimals}f"
    header_html = f"""
    <div style="display:flex;align-items:center;justify-content:space-between;
                background:#0b0e11;padding:14px 18px;border-radius:10px 10px 0 0;
                border:1px solid #1e2329;border-bottom:none;font-family:sans-serif;">
      <div style="display:flex;align-items:center;gap:12px;">
        <span style="background:{badge_color};color:#0b0e11;font-weight:700;
                     padding:4px 12px;border-radius:6px;font-size:14px;">{action}</span>
        <span style="color:#e8e8e8;font-weight:600;font-size:16px;">{symbol_key} · 15m</span>
      </div>
      <div style="color:#8a8f98;font-size:13px;">
        R/R {rr_text} &nbsp;·&nbsp; Confidence {confidence:.0f}% &nbsp;·&nbsp; {ts}
      </div>
    </div>
    """

    if action in ("BUY", "SELL"):
        rows = f"""
        <div style="display:flex;justify-content:space-between;padding:4px 0;color:#ea3943;">
          <span>SL</span><span>{sl:{fmt}}</span>
        </div>"""
        for label, tp, mult in zip(["TP1", "TP2", "TP3"], tps, TP_MULTIPLES):
            rows += f"""
        <div style="display:flex;justify-content:space-between;padding:4px 0;color:#16c784;">
          <span>{label} ({mult}R)</span><span>{tp:{fmt}}</span>
        </div>"""
        levels_html = f"""
        <div style="background:#0b0e11;padding:14px 18px;border:1px solid #1e2329;
                    border-top:none;border-radius:0 0 10px 10px;font-family:sans-serif;
                    font-size:14px;">
          <div style="display:flex;justify-content:space-between;padding:4px 0;color:#f3ba2f;">
            <span>Entry</span><span>{entry:{fmt}}</span>
          </div>
          {rows}
          <div style="color:#5b6270;font-size:11px;margin-top:10px;">
            {engine_label} · Educational analysis only — not financial advice.
          </div>
        </div>
        """
    else:
        levels_html = f"""
        <div style="background:#0b0e11;padding:14px 18px;border:1px solid #1e2329;
                    border-top:none;border-radius:0 0 10px 10px;font-family:sans-serif;
                    font-size:13px;color:#8a8f98;">
          No clear directional edge right now.
          <div style="color:#5b6270;font-size:11px;margin-top:10px;">
            {engine_label} · Educational analysis only — not financial advice.
          </div>
        </div>
        """

    fig = render_chart(feat_df.tail(80), entry, sl, tps, action, decimals)
    return header_html + levels_html, fig


# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------

CUSTOM_CSS = """
body, .gradio-container { background-color: #05070a !important; }
"""

with gr.Blocks(css=CUSTOM_CSS, title="Signal Dashboard (Educational)") as demo:
    gr.Markdown(
        "### 🪙 Signal Dashboard — *Educational, demo-only*\n"
        "XAUUSD uses a PPO model · BTCUSD uses a rule-based technical signal · "
        "TwelveData live feed · Manual execution only, no auto-trading."
    )
    with gr.Row():
        symbol_dd = gr.Dropdown(
            choices=list(SYMBOLS.keys()), value=DEFAULT_SYMBOL,
            label="Symbol", scale=1,
        )
        refresh_btn = gr.Button("🔄 Refresh Signal", variant="primary", scale=2)
    signal_card = gr.HTML()
    chart = gr.Plot()

    refresh_btn.click(fn=run_signal, inputs=[symbol_dd], outputs=[signal_card, chart])
    symbol_dd.change(fn=run_signal, inputs=[symbol_dd], outputs=[signal_card, chart])
    demo.load(fn=run_signal, inputs=[symbol_dd], outputs=[signal_card, chart])

if __name__ == "__main__":
    demo.launch()