Warning: file_exists(): open_basedir restriction in effect. File(/www/wwwroot/com.xiximiao.oa/com.cmstop/public/www/wp-content/db.php) is not within the allowed path(s): (/www/wwwroot/com.xiximiao.oa/com.cmstop/public/www/:/tmp/:/proc/:/var/log/nginx/:/www/wwwroot/com.xiximiao.oa/com.cmstop/public/:/www/wwwroot/com.xiximiao.oa/com.cmstop/vendor/:/www/wwwroot/com.xiximiao.oa/com.cmstop/ppk/) in /www/wwwroot/com.xiximiao.oa/com.cmstop/public/www/wp-includes/load.php on line 707
2341 人工智能强化学习DQN交易智能体策略 » 轻知量化 QMT、PTrade、聚宽策略分享交流平台

2341 人工智能强化学习DQN交易智能体策略

# 标题:人工智能强化学习DQN交易智能体
# 原回测条件:2018-01-01 到 2022-11-01, ¥10000, 每天

from jqdata import *
from jqfactor import *
import numpy as np
import pandas as pd
import pickle
import pandas as pd
import torch
import torch.nn as nn
from tqdm import tqdm
industry_code = ['HY001', 'HY002', 'HY003', 'HY004', 'HY005', 'HY006', 'HY007', 'HY008', 'HY009', 'HY010', 'HY011']


# 初始化函数
def initialize(context):
    # 设定基准
    set_benchmark('000065.XSHE')#000037
    # 用真实价格交易
    set_option('use_real_price', True)
    # 打开防未来函数
    set_option("avoid_future_data", True)
    # 将滑点设置为0
    set_slippage(FixedSlippage(0))
    # 设置交易成本万分之三,不同滑点影响可在归因分析中查看
    set_order_cost(OrderCost(open_tax=0, close_tax=0.001, open_commission=0.0003, close_commission=0.0003,
                             close_today_commission=0, min_commission=5), type='stock')
    # 过滤order中低于error级别的日志
    log.set_level('order', 'error')

    run_daily(adjustment,  '9:30')
    # run_weekly(adjustment, 1, '9:30')

模型加载代码,解锁后查看:

def get_action(context):
    yesterday = context.previous_date
    today = context.current_dt
    initial_list ='000065.XSHE'
    

    df = attribute_history(initial_list, 7, '1d')
    df['涨跌幅'] = df['close'].pct_change() 
    print(df)
    df = df.dropna()

    print(df['涨跌幅'].values)
    df_tensor = torch.Tensor(df['涨跌幅'].values)
    output1 = model_t1(df_tensor)
    action = np.argmax(output1.detach().squeeze(0)) #1:买入 2:卖出 0:持有

    return action
    


# 1-3 整体调整持仓
def adjustment(context):

        initial_list ='000065.XSHE'
        action = get_action(context)
        print(action)
        # 调仓卖出
        value = context.portfolio.cash
        print(value)
        if action == 1:
            log.info("买入")
            open_position(initial_list, value)
        if action == 2:
            log.info("卖出")
            position = context.portfolio.positions[initial_list]
            close_position(position)
        else:
            log.info("持有")


def order_target_value_(security, value):
    if value == 0:
        log.debug("Selling out %s" % (security))
    else:
        log.debug("Order %s to value %f" % (security, value))
    return order_target_value(security, value)


# 3-2 交易模块-开仓
def open_position(security, value):
    order = order_target_value_(security, value)
    if order != None and order.filled > 0:
        return True
    return False


# 3-3 交易模块-平仓
def close_position(position):
    security = position.security
    order = order_target_value_(security, 0)  # 可能会因停牌失败
    if order != None:
        if order.status == OrderStatus.held and order.filled == order.amount:
            return True
    return False


# 4-2 清仓后次日资金可转
def close_account(context):

        if len(g.hold_list) != 0:
            for stock in g.hold_list:
                position = context.portfolio.positions[stock]
                close_position(position)
                log.info("卖出[%s]" % (stock))


 

配套模型训练代码: 在研究环境中新建

import pandas as pd
import os
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
import matplotlib.pyplot as plt
from tqdm import tqdm

class stock:
    def __init__(self, df,  window_size=6):
        self.n_actions = 3  # 动作数量
        self.n_features = window_size  # 特征数量
        self.trend = df['close'].values  # 收盘数据
        self.trend_open = df['open'].values
        self.window_size = window_size  # 滑动窗口大小
        self.half_window = window_size // 2
        self.hold_num = 0

    def step(self, action):

        if action == 1:
            self.hold_num = 1
        if action == 2:
            self.hold_num = 0
        if action == 0:
            self.hold_num = self.hold_num

        # self.reward = (self.maket_value - self.last_value) / self.last_value
        reward = (self.trend[self.t + 1] - self.trend[self.t]) / self.trend[self.t]
        #         reward = (self.trend_open[self.t + 2] - self.trend_open[self.t + 1]) / self.trend_open[self.t + 1]
        if np.abs(reward) <= 0.015:
            self.reward = reward * 0.2
        elif np.abs(reward) <= 0.03:
            self.reward = reward * 0.7
        elif np.abs(reward) >= 0.05:
            if reward < 0:
                self.reward = (reward + 0.05) * 0.1 - 0.05
            else:
                self.reward = (reward - 0.05) * 0.1 + 0.05

        # reward = (self.trend[self.t + 1] - self.trend[self.t]) / self.trend[self.t]
        if self.hold_num > 0 or action == 2:
            self.reward = reward
            if action == 2:
                self.reward = -self.reward
        else:
            self.reward = -self.reward * 0.1
            # self.reward = 0

        done = False
        self.t = self.t + 1
        if self.t == len(self.trend) - 2:
            done = True
        s_ = self.get_state(self.t)
        reward = self.reward
#         print(reward)
        return s_, reward, done

    def get_state(self, t):  # 某t时刻的状态
        window_size = self.window_size + 1
        d = t - window_size + 1
        block = []
        if d < 0:
            for i in range(-d):
                block.append(self.trend[0])
            for i in range(t + 1):
                block.append(self.trend[i])
        else:
            block = self.trend[d: t + 1]

        res = []
        for i in range(window_size - 1):
            res.append((block[i + 1] - block[i]) / (block[i] + 0.0001))  # 每步收益
        return np.array(res)  # 作为状态编码

    def reset(self):
        self.total_profit = 0  # 总盈利
        self.t = self.window_size // 2  # 时间
        self.reward = 0  # 收益
        return self.get_state(self.t)


class DQN(nn.Module):
    def __init__(self, input_shape, n_actions):
        super(DQN, self).__init__()
        units = 32
        self.fc1 = nn.Linear(input_shape, units)
        self.relu = nn.ReLU()
        self.fc2 = nn.Linear(units, n_actions)

    def forward(self, x):
        x = self.fc1(x)
        x = self.relu(x)
        x = self.fc2(x)
        return x




os.environ['KMP_DUPLICATE_LIB_OK'] = 'True'
device = 'cpu'
np.random.seed(1)
torch.manual_seed(41)

date='2022-11-01'
by_date = '2018-01-01'
df= get_price('000065.XSHE', 
              start_date=by_date,
              end_date=date,
             frequency='1d', )
print(df.head(7))
print(df.tail(7))
env = stock(df)



max_round = 100
step = 0
LOSS = []
REWORD = []
net = DQN(6, 3)
net.train()
tgt_net = DQN(6, 3)
tgt_net.train()
learn_step_counter = 0
replace_target_iter = 200
batch_size = 512
lr = 0.001
gamma = 0.9

epsilon = 200
epsilon_increment = None
epsilon_max = 0.9
memory_size = 4000
n_features = 6
optimizer = optim.Adam(net.parameters(), lr=lr)
min_validation_loss = 9999
for episode in tqdm(range(max_round)):

    # initial observation
    observation = env.reset()
    l = 0
    r = 0
    memory_counter = 0
    memory = np.zeros((memory_size, n_features * 2 + 2))
    while True:
        Observation = [observation[np.newaxis, :]]
        Observation = torch.tensor(Observation, dtype=torch.float32).to(device)
        # forward feed the observation and get q value for every actions
        actions_value = net(Observation).detach().cpu().squeeze(0)
        action = np.argmax(actions_value)
#         print(action)
        # RL take action and get next observation and reward
        observation_, reward, done = env.step(action)
        r = r + reward
        transition = np.hstack((observation, [action, reward], observation_))
        # replace the old memory with new memory
        index = memory_counter % memory_size
        memory[index, :] = transition
        memory_counter += 1
        if learn_step_counter % replace_target_iter == 0:
            tgt_net.load_state_dict(net.state_dict())
        # sample batch memory from all memory
        if memory_counter > memory_size:
            sample_index = np.random.choice(memory_size, size=batch_size)
        else:
            sample_index = np.random.choice(memory_counter, size=batch_size)
        batch_memory = memory[sample_index, :]

        s_ = torch.tensor(batch_memory[:, -n_features:], dtype=torch.float32).to(device)
        s = torch.tensor(batch_memory[:, :n_features], dtype=torch.float32).to(device)
        eval_act_index = batch_memory[:, n_features].astype(int)
        reward = torch.tensor(batch_memory[:, n_features + 1], dtype=torch.float32)
        q_next = tgt_net(s_)
        q_eval = net(s)
        # # change q_target w.r.t q_eval's action
        q_target = q_eval.clone()
        batch_index = np.arange(batch_size, dtype=np.int32)
        max_, _ = torch.max(q_next, dim=1)

        q_target[batch_index, eval_act_index] = reward + gamma * max_

        # loss backpropogation
        optimizer.zero_grad()
        loss = nn.MSELoss()(q_eval, q_target)

        loss.backward()
        optimizer.step()


        # increasing epsilon
        epsilon = epsilon + epsilon_increment if epsilon < epsilon_max else epsilon_max
        learn_step_counter += 1

        # swap observation
        observation = observation_

        # break while loop when end of this episode
        if done:
            break
        step += 1
        l = l + loss.detach().cpu().numpy()
    if min_validation_loss > l:
        min_validation_loss = l
        best_epoch = episode
        print('Min loss ' + str(min_validation_loss) + ' in epoch ' + str(best_epoch))
        torch.save(net.state_dict(), 'tgt_net.pt')

    LOSS.append(l)
    REWORD.append(r)
 
x_values = np.arange(len(LOSS))  # 假设epoch是从0开始的整数序列
plt.plot(x_values, LOSS)
plt.xlabel('Epoch')
plt.ylabel('Value (y)')
plt.title('Plotting y values over epochs')
plt.show()

plt.plot(x_values, REWORD)
plt.xlabel('Epoch')
plt.ylabel('Value (y)')
plt.title('Plotting y values over epochs')
plt.show()
2025-02-22
⚠️
本站资源大多来自网络,仅供网友学习交流,未经作者或上传书面授权,请勿作他用。
站长 vx: xiangyin615 或者 留言反馈 ,我们将尽快处理。
Notice: When you of the legal rights be violate, please stir to vx: xiangyin615
个人中心
购物车
优惠劵
搜索