import os
import sys
import random
import numpy as np
import pandas as pd
import torch
import torch.nn.functional as F
import dolphindb as ddb
from torch.utils.data import Dataset, DataLoader
import time

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.dirname(BASE_DIR)
sys.path.append(PROJECT_ROOT)

from model import Kronos, KronosTokenizer


# ============================================================
# 1. basic config
# ============================================================
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.dirname(BASE_DIR)

DDB_HOST = '192.168.100.43'
DDB_PORT = 8742
DDB_USER = 'admin'
DDB_PASSWORD = 'DolphinDB123456'

DDB_TABLE = 'loadTable("dfs://tushare_factor_minute_db", "factor_minute_1min")'
SYMBOL = '000514.SZ'

START_TIME = pd.Timestamp('2025-01-01 00:00:00')
END_TIME = pd.Timestamp('2025-07-01 00:00:00')

LOOKBACK_WINDOW = 512
PREDICT_WINDOW = 48
WINDOW = LOOKBACK_WINDOW + PREDICT_WINDOW + 1
CLIP = 5.0

TRAIN_RATIO = 0.8
VAL_RATIO = 0.2

BATCH_SIZE = 32
NUM_WORKERS = 4
SEED = 10

TOKENIZER_EPOCHS = 10
PREDICTOR_EPOCHS = 10
TOKENIZER_LR = 2e-4
PREDICTOR_LR = 4e-5

SAVE_DIR = os.path.join(PROJECT_ROOT, 'outputs', 'ddb_factor_train_single_stock')
TOKENIZER_SAVE_DIR = os.path.join(SAVE_DIR, 'tokenizer', 'best_model')
PREDICTOR_SAVE_DIR = os.path.join(SAVE_DIR, 'predictor', 'best_model')
PRETRAINED_PREDICTOR_PATH = 'NeoQuasar/Kronos-small'


# ============================================================
# 2. helper
# ============================================================

def set_seed(seed):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


class FactorDataset(Dataset):
    def __init__(self, df):
        self.factor_array = df[FACTOR_COLUMNS].values.astype(np.float32)
        self.stamp_array = df[['minute', 'hour', 'weekday', 'day', 'month']].values.astype(np.float32)
        self.n_samples = len(df) - WINDOW + 1

    def __len__(self):
        return self.n_samples

    def __getitem__(self, idx):
        x = self.factor_array[idx:idx + WINDOW].copy()
        x_stamp = self.stamp_array[idx:idx + WINDOW].copy()

        x_mean = np.mean(x, axis=0)
        x_std = np.std(x, axis=0)
        x = (x - x_mean) / (x_std + 1e-5)
        x = np.clip(x, -CLIP, CLIP)

        return torch.from_numpy(x), torch.from_numpy(x_stamp)


# ============================================================
# 3. load factor data from ddb, single stock only
# ============================================================

print('load factor data from ddb...')
s = ddb.session()
s.connect(DDB_HOST, DDB_PORT, DDB_USER, DDB_PASSWORD)

batch_list = []
current_time = START_TIME

while current_time < END_TIME:
    script = f'''
        select factorvalue
        from {DDB_TABLE}
        where trade_date = {current_time.strftime('%Y.%m.%d')} and code = `{SYMBOL}
        pivot by concatDateTime(trade_date,trade_time) as timestamps, factorname
    '''
    batch_df = s.run(script)
    if len(batch_df) > 0:
        batch_list.append(batch_df)
        print(current_time.strftime('%Y-%m-%d'), len(batch_df))
    current_time = current_time + pd.Timedelta(days=1)

s.close()

df = pd.concat(batch_list, ignore_index=True)
FACTOR_COLUMNS = df.columns.to_list()
del FACTOR_COLUMNS[0]


# ============================================================
# 4. factor preprocess
# ============================================================

print('preprocess factor data...')
df = df.sort_values('timestamps').reset_index(drop=True)
df[FACTOR_COLUMNS] = df[FACTOR_COLUMNS].ffill().fillna(0.0)

df['minute'] = df['timestamps'].dt.minute
df['hour'] = df['timestamps'].dt.hour
df['weekday'] = df['timestamps'].dt.weekday
df['day'] = df['timestamps'].dt.day
df['month'] = df['timestamps'].dt.month

print('data range:', df['timestamps'].min(), '->', df['timestamps'].max())
print('total rows:', len(df))
print('factor dim:', len(FACTOR_COLUMNS))


# ============================================================
# 5. train / val split
# ============================================================

train_end = int(len(df) * TRAIN_RATIO)
val_end = int(len(df) * (TRAIN_RATIO + VAL_RATIO))

train_df = df.iloc[:train_end].copy()
val_df = df.iloc[train_end:val_end].copy()

print('train rows:', len(train_df))
print('val rows:', len(val_df))

train_dataset = FactorDataset(train_df)
val_dataset = FactorDataset(val_df)

train_loader = DataLoader(
    train_dataset,
    batch_size=BATCH_SIZE,
    shuffle=True,
    num_workers=NUM_WORKERS,
    pin_memory=True,
    drop_last=True,
)

val_loader = DataLoader(
    val_dataset,
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=NUM_WORKERS,
    pin_memory=True,
    drop_last=False,
)


# ============================================================
# 6. tokenizer training
# ============================================================

set_seed(SEED)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
os.makedirs(TOKENIZER_SAVE_DIR, exist_ok=True)
os.makedirs(PREDICTOR_SAVE_DIR, exist_ok=True)

print('device:', device)
print('start tokenizer training...')

tokenizer = KronosTokenizer(
    d_in=len(FACTOR_COLUMNS),
    d_model=256,
    n_heads=4,
    ff_dim=512,
    n_enc_layers=4,
    n_dec_layers=4,
    ffn_dropout_p=0.0,
    attn_dropout_p=0.0,
    resid_dropout_p=0.0,
    s1_bits=10,
    s2_bits=10,
    beta=0.05,
    gamma0=1.0,
    gamma=1.1,
    zeta=0.05,
    group_size=4,
).to(device)

tokenizer_optimizer = torch.optim.AdamW(
    tokenizer.parameters(),
    lr=TOKENIZER_LR,
    weight_decay=0.1,
)

tokenizer_scheduler = torch.optim.lr_scheduler.OneCycleLR(
    tokenizer_optimizer,
    max_lr=TOKENIZER_LR,
    steps_per_epoch=len(train_loader),
    epochs=TOKENIZER_EPOCHS,
    pct_start=0.03,
    div_factor=10,
)

best_tokenizer_val_loss = float('inf')

for epoch in range(TOKENIZER_EPOCHS):
    tokenizer.train()
    train_loss_sum = 0.0

    for batch_x, _ in train_loader:
        batch_x = batch_x.to(device, non_blocking=True)

        zs, bsq_loss, _, _ = tokenizer(batch_x)
        z_pre, z = zs

        recon_loss_pre = F.mse_loss(z_pre, batch_x)
        recon_loss_all = F.mse_loss(z, batch_x)
        recon_loss = recon_loss_pre + recon_loss_all
        loss = (recon_loss + bsq_loss) / 2

        tokenizer_optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(tokenizer.parameters(), max_norm=2.0)
        tokenizer_optimizer.step()
        tokenizer_scheduler.step()

        train_loss_sum += loss.item()

    tokenizer.eval()
    val_loss_sum = 0.0

    with torch.no_grad():
        for batch_x, _ in val_loader:
            batch_x = batch_x.to(device, non_blocking=True)
            zs, _, _, _ = tokenizer(batch_x)
            _, z = zs
            val_loss_sum += F.mse_loss(z, batch_x).item()

    avg_train_loss = train_loss_sum / len(train_loader)
    avg_val_loss = val_loss_sum / len(val_loader)
    print(f'tokenizer epoch {epoch + 1}/{TOKENIZER_EPOCHS} train_loss={avg_train_loss:.6f} val_loss={avg_val_loss:.6f}')

    if avg_val_loss < best_tokenizer_val_loss:
        best_tokenizer_val_loss = avg_val_loss
        tokenizer.save_pretrained(TOKENIZER_SAVE_DIR)
        print('save tokenizer:', TOKENIZER_SAVE_DIR)


# ============================================================
# 7. predictor training
# ============================================================

print('start predictor training...')
tokenizer = KronosTokenizer.from_pretrained(TOKENIZER_SAVE_DIR).to(device)
tokenizer.eval()

predictor = Kronos.from_pretrained(PRETRAINED_PREDICTOR_PATH).to(device)
predictor_optimizer = torch.optim.AdamW(
    predictor.parameters(),
    lr=PREDICTOR_LR,
    betas=(0.9, 0.95),
    weight_decay=0.1,
)

predictor_scheduler = torch.optim.lr_scheduler.OneCycleLR(
    predictor_optimizer,
    max_lr=PREDICTOR_LR,
    steps_per_epoch=len(train_loader),
    epochs=PREDICTOR_EPOCHS,
    pct_start=0.03,
    div_factor=10,
)

best_predictor_val_loss = float('inf')

for epoch in range(PREDICTOR_EPOCHS):
    predictor.train()
    train_loss_sum = 0.0

    for batch_x, batch_x_stamp in train_loader:
        batch_x = batch_x.to(device, non_blocking=True)
        batch_x_stamp = batch_x_stamp.to(device, non_blocking=True)

        with torch.no_grad():
            token_seq_0, token_seq_1 = tokenizer.encode(batch_x, half=True)

        token_in = [token_seq_0[:, :-1], token_seq_1[:, :-1]]
        token_out = [token_seq_0[:, 1:], token_seq_1[:, 1:]]

        logits = predictor(token_in[0], token_in[1], batch_x_stamp[:, :-1, :])
        loss, s1_loss, s2_loss = predictor.head.compute_loss(
            logits[0], logits[1], token_out[0], token_out[1]
        )

        predictor_optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(predictor.parameters(), max_norm=3.0)
        predictor_optimizer.step()
        predictor_scheduler.step()

        train_loss_sum += loss.item()

    predictor.eval()
    val_loss_sum = 0.0

    with torch.no_grad():
        for batch_x, batch_x_stamp in val_loader:
            batch_x = batch_x.to(device, non_blocking=True)
            batch_x_stamp = batch_x_stamp.to(device, non_blocking=True)

            token_seq_0, token_seq_1 = tokenizer.encode(batch_x, half=True)
            token_in = [token_seq_0[:, :-1], token_seq_1[:, :-1]]
            token_out = [token_seq_0[:, 1:], token_seq_1[:, 1:]]

            logits = predictor(token_in[0], token_in[1], batch_x_stamp[:, :-1, :])
            loss, _, _ = predictor.head.compute_loss(
                logits[0], logits[1], token_out[0], token_out[1]
            )
            val_loss_sum += loss.item()

    avg_train_loss = train_loss_sum / len(train_loader)
    avg_val_loss = val_loss_sum / len(val_loader)
    print(f'predictor epoch {epoch + 1}/{PREDICTOR_EPOCHS} train_loss={avg_train_loss:.6f} val_loss={avg_val_loss:.6f}')

    if avg_val_loss < best_predictor_val_loss:
        best_predictor_val_loss = avg_val_loss
        predictor.save_pretrained(PREDICTOR_SAVE_DIR)
        print('save predictor:', PREDICTOR_SAVE_DIR)


# ============================================================
# 8. encode / decode demo
# ============================================================

print('run encode / decode demo...')
sample_x, _ = next(iter(val_loader))
sample_x = sample_x[:1].to(device)

with torch.no_grad():
    sample_tokens = tokenizer.encode(sample_x, half=True)
    sample_recon = tokenizer.decode(sample_tokens, half=True)
    sample_mse = F.mse_loss(sample_recon, sample_x).item()

print('sample token shape s1:', sample_tokens[0].shape)
print('sample token shape s2:', sample_tokens[1].shape)
print('sample recon shape:', sample_recon.shape)
print('sample recon mse:', sample_mse)

time.sleep(1)