Kronos 项目功能调用流程文档

Kronos 项目功能调用流程文档

目录

  1. 项目概述
  2. 预测功能调用流程
  3. 批量预测功能流程
  4. 模型加载流程
  5. 微调训练流程
  6. Web界面功能流程
  7. 核心组件交互关系

项目概述

Kronos项目基于分层Tokenization和Transformer架构,为金融时间序列预测提供完整的解决方案。本文档详细描述了各项功能的调用流程和组件间的交互关系。

核心架构组件

KronosPredictor (预测器接口)
    ├── KronosTokenizer (数据分词器)
    ├── Kronos (主Transformer模型)
    └── auto_regressive_inference (自回归推理)

预测功能调用流程

1. 基础预测调用

入口文件: examples/prediction_example.py

加载模型和分词器

初始化预测器

准备输入数据

执行预测

结果后处理

可视化输出

详细调用序列

步骤1: 模型加载 (prediction_example.py:42-43)

# 从Hugging Face Hub加载预训练模型
tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base")
model = Kronos.from_pretrained("NeoQuasar/Kronos-small")

步骤2: 预测器初始化 (prediction_example.py:46)

predictor = KronosPredictor(model, tokenizer, device="cuda:0", max_context=512)

步骤3: 数据准备 (prediction_example.py:48-57)

# 加载CSV数据
df = pd.read_csv("./data/XSHG_5min_600977.csv")
df['timestamps'] = pd.to_datetime(df['timestamps'])

# 定义上下文窗口和预测长度
lookback = 400
pred_len = 120

# 提取特征数据
x_df = df.loc[:lookback-1, ['open', 'high', 'low', 'close', 'volume', 'amount']]
x_timestamp = df.loc[:lookback-1, 'timestamps']
y_timestamp = df.loc[lookback:lookback+pred_len-1, 'timestamps']

步骤4: 执行预测 (prediction_example.py:60-69)

pred_df = predictor.predict(
    df=x_df,
    x_timestamp=x_timestamp,
    y_timestamp=y_timestamp,
    pred_len=pred_len,
    T=1.0,          # 采样温度
    top_p=0.9,      # 核采样概率
    sample_count=1, # 采样次数
    verbose=True
)

2. 内部详细调用流程

KronosPredictor.predict() 方法流程

文件: model/kronos.py:483-523

输入验证

数据预处理

时间特征提取

数据归一化

张量转换

生成预测

反归一化

返回DataFrame

详细步骤:

  1. 输入验证 (kronos.py:485-500)

    # 检查DataFrame格式
    # 验证必需列: ['open', 'high', 'low', 'close']
    # 处理缺失的volume/amount列
    # 检查NaN值
    
  2. 时间特征提取 (kronos.py:501-506)

    # 调用 calc_time_stamps(x_timestamp)
    # 提取: minute, hour, weekday, day, month
    
  3. 数据归一化 (kronos.py:508-511)

    # 计算均值和标准差
    # 归一化: (x - x_mean) / (x_std + 1e-5)
    # 裁剪值: np.clip(x, -self.clip, self.clip)
    
  4. 生成预测 (kronos.py:517)

    # 调用 self.generate() 方法
    preds = self.generate(x, x_stamp, y_stamp, pred_len, T, top_k, top_p, sample_count, verbose)
    
auto_regressive_inference() 核心推理流程

文件: model/kronos.py:389-443

输入准备

数据分词

自回归循环

S1令牌预测

S2令牌预测

令牌更新

是否完成?

解码输出

后处理

关键步骤:

  1. 分词处理 (kronos.py:400)

    # 使用 tokenizer.encode(x, half=True) 将连续数据转换为离散令牌
    token_in = tokenizer.encode(x, half=True)
    
  2. 自回归循环 (kronos.py:414-433)

    for i in range(pred_len):
        # 动态上下文管理
        x_stamp_dynamic = get_dynamic_stamp(x_stamp[:, i:, :], y_stamp[:, :i+1, :], max_context)
    
        # S1令牌预测
        s1_logits = model.decode_s1(token_in, x_stamp_dynamic)
        s1_tokens = sample_from_logits(s1_logits, T, top_k, top_p)
    
        # S2令牌预测 (基于S1条件)
        s2_logits = model.decode_s2(token_in, s1_tokens, x_stamp_dynamic)
        s2_tokens = sample_from_logits(s2_logits, T, top_k, top_p)
    
        # 更新令牌序列
        token_in[0] = torch.cat([token_in[0], s1_tokens], dim=1)
        token_in[1] = torch.cat([token_in[1], s2_tokens], dim=1)
    
  3. 解码输出 (kronos.py:437-441)

    # 使用 tokenizer.decode() 将令牌转换回数据空间
    preds = tokenizer.decode(token_in[0], token_in[1])
    

批量预测功能流程

1. 批量预测调用

入口文件: examples/prediction_batch_example.py

加载模型

准备多个数据序列

批量数据验证

批量预处理

并行推理

单独反归一化

返回结果列表

调用示例

步骤1: 准备多个数据序列 (prediction_batch_example.py:55-65)

dfs = []
xtsp = []
ytsp = []

for i in range(5):
    # 为每个序列提取不同的时间段
    idf = df.loc[(i*400):(i*400+lookback-1), ['open', 'high', 'low', 'close', 'volume', 'amount']]
    i_x_timestamp = df.loc[(i*400):(i*400+lookback-1), 'timestamps']
    i_y_timestamp = df.loc[(i*400+lookback):(i*400+lookback+pred_len-1), 'timestamps']

    dfs.append(idf)
    xtsp.append(i_x_timestamp)
    ytsp.append(i_y_timestamp)

步骤2: 执行批量预测 (prediction_batch_example.py:67-72)

pred_df = predictor.predict_batch(
    df_list=dfs,
    x_timestamp_list=xtsp,
    y_timestamp_list=ytsp,
    pred_len=pred_len,
)

2. 内部批量处理流程

KronosPredictor.predict_batch() 方法

文件: model/kronos.py:526-625

处理步骤:

  1. 输入验证 (kronos.py:546-576)

    # 检查列表类型
    # 验证列表长度一致性
    # 检查每个DataFrame的必需列
    # 验证时间戳长度
    
  2. 单个序列处理 (kronos.py:561-604)

    for i, (df, x_timestamp, y_timestamp) in enumerate(zip(df_list, x_timestamp_list, y_timestamp_list)):
        # 数据验证
        # 时间特征提取
        # 数据提取和归一化
        # 存储处理后的数据、均值、标准差
    
  3. 批量张量创建 (kronos.py:612-614)

    # 堆叠数据为批量张量
    x_batch = torch.tensor(x_batch, dtype=torch.float32, device=self.device)
    x_stamp_batch = torch.tensor(x_stamp_batch, dtype=torch.float32, device=self.device)
    y_stamp_batch = torch.tensor(y_stamp_batch, dtype=torch.float32, device=self.device)
    
  4. 批量生成 (kronos.py:616)

    # 调用 generate() 方法处理整个批次
    preds = self.generate(x_batch, x_stamp_batch, y_stamp_batch, pred_len, T, top_k, top_p, sample_count, verbose)
    
  5. 单独反归一化 (kronos.py:619-625)

    # 对每个序列应用单独的均值/标准差
    pred_list = []
    for i in range(len(df_list)):
        pred_i = preds[i] * (x_stds[i] + 1e-5) + x_means[i]
        pred_list.append(pd.DataFrame(pred_i.cpu().numpy(), index=y_timestamp_list[i], columns=self.columns))
    

模型加载流程

1. Hugging Face Hub模型加载

KronosTokenizer加载流程

文件: model/kronos.py:13-179

PyTorchModelHubMixin

下载模型文件

初始化嵌入层

创建Transformer块

初始化BSQuantizer

加载权重

返回分词器实例

关键初始化步骤:

# 嵌入层初始化
self.embed = nn.Linear(self.d_in, self.d_model)

# 创建编码器和解码器Transformer块
self.encoder_blocks = nn.ModuleList([...])
self.decoder_blocks = nn.ModuleList([...])

# 初始化BSQuantizer
self.tokenizer = BSQuantizer(...)
Kronos模型加载流程

文件: model/kronos.py:180-329

PyTorchModelHubMixin

下载模型权重

初始化分层嵌入

创建时间嵌入

构建Transformer块

初始化双头输出

加载预训练权重

设置评估模式

关键组件初始化:

# 分层嵌入 (S1, S2令牌)
self.embedding = HierarchicalEmbedding(d_model, n_heads)

# 时间嵌入
self.temporal_embedding = TemporalEmbedding(d_model)

# Transformer块列表
self.layers = nn.ModuleList([...])

# 双头输出
self.head = DualHead(d_model, vocab_size)

2. Web界面模型加载

文件: webui/app.py:626-663

接收模型加载请求

选择模型配置

获取模型ID

加载分词器

加载主模型

创建预测器

缓存模型实例

返回加载状态

API调用示例:

@app.route('/api/load-model', methods=['POST'])
def load_model():
    # 1. 选择模型配置
    model_config = AVAILABLE_MODELS[request.json.get('model')]

    # 2. 加载模型
    tokenizer = KronosTokenizer.from_pretrained(model_config['tokenizer_id'])
    model = Kronos.from_pretrained(model_config['model_id'])

    # 3. 创建预测器
    predictor = KronosPredictor(
        model=model,
        tokenizer=tokenizer,
        device=device,
        max_context=model_config['context_length']
    )

微调训练流程

1. 完整微调管道

配置设置

数据预处理

分词器微调

预测器微调

回测评估

性能分析

2. 分词器微调流程

入口: finetune/train_tokenizer.py

调用命令:

torchrun --standalone --nproc_per_node=NUM_GPUS finetune/train_tokenizer.py

训练流程:

  1. DDP初始化 - 设置分布式训练环境
  2. 数据加载 - 加载预处理的Qlib数据
  3. 优化器设置 - 配置AdamW优化器和学习率调度
  4. 训练循环 - 分词器特定的训练逻辑
  5. 模型保存 - 保存最佳分词器检查点

3. 预测器微调流程

文件: finetune/train_predictor.py

主要训练流程

函数调用链:

main()
├── setup_ddp()                    # DDP环境设置
├── create_dataloaders()           # 创建数据加载器
└── train_model()                  # 主训练函数
    ├── train_epoch()             # 训练轮次
    │   ├── tokenizer.encode()    # 数据分词
    │   ├── model()               # 前向传播
    │   └── loss.backward()       # 反向传播
    └── validate_epoch()          # 验证轮次
详细训练步骤

1. 数据加载器创建 (train_predictor.py:29-56)

def create_dataloaders(config: dict, rank: int, world_size: int):
    # 创建分布式数据集
    train_dataset = QlibDataset('train')
    valid_dataset = QlibDataset('val')

    # 创建分布式采样器
    train_sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=rank, shuffle=True)
    val_sampler = DistributedSampler(valid_dataset, num_replicas=world_size, rank=rank, shuffle=False)

    # 创建数据加载器
    train_loader = DataLoader(train_dataset, batch_size=config['batch_size'], sampler=train_sampler)
    val_loader = DataLoader(valid_dataset, batch_size=config['batch_size'], sampler=val_sampler)

2. 模型训练循环 (train_predictor.py:87-176)

for epoch in range(config['epochs']):
    # 训练阶段
    model.train()
    for batch_idx, batch in enumerate(train_loader):
        # 数据分词
        token_in = tokenizer.encode(batch_x, half=True)

        # 前向传播
        output = model(token_in[0], token_in[1], batch_x_stamp[:, :-1, :])

        # 损失计算
        loss = model.module.head.compute_loss(output, target)

        # 反向传播
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), config['grad_clip'])
        optimizer.step()

    # 验证阶段
    model.eval()
    with torch.no_grad():
        for batch in val_loader:
            # 验证逻辑...

    # 保存最佳模型
    if val_loss < best_val_loss:
        torch.save(model.module.state_dict(), best_model_path)

4. 数据集处理流程

文件: finetune/dataset.py

QlibDataset初始化流程

加载配置

加载pickle数据

索引预计算

时间特征提取

随机采样设置

数据集就绪

关键处理步骤 (dataset.py:23-86):

def __init__(self, split='train'):
    # 1. 配置加载
    self.config = Config()

    # 2. 数据加载
    data = pd.read_pickle(f"{self.config.dataset_path}/{split}_data.pkl")

    # 3. 索引预计算
    for symbol in self.data.keys():
        # 为每个符号生成时间特征
        time_features = calc_time_stamps(data[symbol]['timestamps'])
        # 预计算所有有效的滑动窗口索引
        indices = [(symbol, i) for i in range(len(data[symbol]) - context_len - pred_len)]
        self.indices.extend(indices)

    # 4. 随机采样设置
    self.rng = np.random.default_rng(epoch_seed)

Web界面功能流程

1. Flask应用架构

文件: webui/app.py

Flask应用初始化

模型加载管理

文件上传处理

预测API

结果可视化

前端展示

2. 主要API端点

模型加载API (/api/load-model)

文件: webui/app.py:626-663

@app.route('/api/load-model', methods=['POST'])
def load_model():
    # 1. 获取请求参数
    model_name = request.json.get('model')
    device = request.json.get('device', 'cpu')

    # 2. 模型配置选择
    model_config = AVAILABLE_MODELS[model_name]

    # 3. 模型加载
    tokenizer = KronosTokenizer.from_pretrained(model_config['tokenizer_id'])
    model = Kronos.from_pretrained(model_config['model_id'])

    # 4. 预测器创建
    predictor = KronosPredictor(
        model=model,
        tokenizer=tokenizer,
        device=device,
        max_context=model_config['context_length']
    )

    # 5. 全局变量存储
    global tokenizer, model, predictor
    tokenizer, model, predictor = tokenizer, model, predictor

    return jsonify({'status': 'success', 'model': model_config})
预测API (/api/predict)

文件: webui/app.py:404-624

处理流程:

接收预测请求

验证模型状态

文件上传处理

数据预处理

时间周期处理

执行预测

结果处理

图表生成

返回JSON响应

关键处理步骤 (app.py:421-490):

  1. 文件处理 (app.py:439-453)

    # 保存上传的文件
    file = request.files['file']
    filename = secure_filename(file.filename)
    file_path = os.path.join(UPLOAD_FOLDER, filename)
    file.save(file_path)
    
    # 加载数据
    df = pd.read_csv(file_path)
    df['timestamps'] = pd.to_datetime(df['timestamps'])
    
  2. 时间周期处理 (app.py:467-478)

    if start_date:
        # 自定义时间周期
        start_timestamp = pd.to_datetime(start_date)
        data_end_idx = df[df['timestamps'] >= start_timestamp].index[0]
    
        # 提取历史和预测数据
        x_df = df.loc[data_end_idx-lookback:data_end_idx-1, required_columns]
        x_timestamp = df.loc[data_end_idx-lookback:data_end_idx-1, 'timestamps']
        y_timestamp = generate_future_timestamps(x_timestamp.iloc[-1], pred_len)
    else:
        # 使用最新数据
        x_df = df.tail(lookback)[required_columns]
        x_timestamp = df.tail(lookback)['timestamps']
        y_timestamp = generate_future_timestamps(x_timestamp.iloc[-1], pred_len)
    
  3. 预测执行 (app.py:479-490)

    try:
        pred_df = predictor.predict(
            df=x_df,
            x_timestamp=x_timestamp,
            y_timestamp=y_timestamp,
            pred_len=pred_len,
            T=temperature,
            top_p=top_p,
            sample_count=sample_count,
            verbose=False
        )
    except Exception as e:
        return jsonify({'error': str(e)}), 500
    
  4. 结果处理和可视化 (app.py:494-612)

    # 实际数据提取
    actual_df = df.loc[data_end_idx:data_end_idx+pred_len-1, required_columns]
    
    # 创建预测图表
    chart_json = create_prediction_chart(
        historical_data=df.tail(lookback*2),
        actual_data=actual_df,
        predicted_data=pred_df,
        prediction_start=prediction_start_timestamp
    )
    
    # 保存预测结果
    result_data = save_prediction_results(
        actual_data=actual_df,
        predicted_data=pred_df,
        model_info=model_info
    )
    

核心组件交互关系

1. 模型组件层次结构

KronosPredictor

KronosTokenizer

Kronos Model

Sampling Functions

BSQuantizer

Encoder/Decoder Blocks

HierarchicalEmbedding

TemporalEmbedding

Transformer Blocks

DualHead

top_k_top_p_filtering

sample_from_logits

2. 数据流转图

Raw CSV Data

KronosPredictor

Data Validation

Time Feature Extraction

Normalization

KronosTokenizer

Token Indices

Kronos Model

Auto-regressive Inference

Token Decoding

Denormalization

Prediction DataFrame

3. 训练时组件交互

QlibDataset

DataLoader

Training Loop

KronosTokenizer.encode

Kronos Model Forward

DualHead Loss

Backpropagation

Optimizer Step

Model Checkpoint

4. 推理时组件交互

No

Yes

Input DataFrame

KronosPredictor.predict

Preprocessing Pipeline

auto_regressive_inference

S1 Token Prediction

S2 Token Prediction

Token Update

Prediction Complete?

Tokenizer.decode

Post-processing

Output DataFrame


性能优化要点

1. 批量处理优化

  • GPU并行利用:批量预测充分利用GPU并行计算能力
  • 内存管理:避免不必要的数据拷贝和转换
  • 上下文缓存:复用时间特征计算结果

2. 推理优化

  • 自回归效率:动态上下文窗口管理
  • 采样优化:核采样和温度采样的高效实现
  • 张量操作:最小化CPU-GPU数据传输

3. 训练优化

  • 分布式训练:DDP支持多GPU并行训练
  • 梯度裁剪:防止梯度爆炸
  • 学习率调度:OneCycleLR优化训练收敛

错误处理和异常情况

1. 数据验证错误

# 缺少必需列
if not all(col in df.columns for col in ['open', 'high', 'low', 'close']):
    raise ValueError("Missing required columns")

# 数据长度不足
if len(df) < lookback + pred_len:
    raise ValueError("Insufficient data length")

2. 模型加载错误

# Hugging Face Hub连接失败
try:
    model = Kronos.from_pretrained(model_id)
except Exception as e:
    raise ModelLoadError(f"Failed to load model: {e}")

3. 推理错误

# GPU内存不足
if torch.cuda.is_available():
    torch.cuda.empty_cache()
else:
    raise RuntimeError("CUDA not available")

总结

Kronos项目的功能调用流程体现了以下设计特点:

  1. 模块化设计:清晰的组件分层和职责分离
  2. 灵活的接口:支持单序列和批量预测
  3. 完整的训练管道:从数据准备到模型评估的端到端流程
  4. 用户友好的Web界面:简化模型使用和结果可视化
  5. 性能优化:充分利用GPU并行和分布式训练能力

通过理解这些调用流程,开发者可以更好地使用、扩展和优化Kronos项目的各项功能。

股票5、15、30、60分钟以及日线数据下载工具 (CSV格式)功能说明这个工具用于从baostock下载股票5、15、30、60分钟以及日线的K线数据,并保存为CSV格式,兼容Kronos项目的数据格式要求。主要特性 灵活的股票选择: 支持指定单只或多只股票,也支持下载沪深300全部股票 CSV格式存储: 保存为标准CSV格式,便于后续处理和分析 自定义日期范围: 支持指定任意日期范围进行数据下载 防频率限制: 内置延时机制,避免请求过于频繁 格式兼容: 完全兼容现有Kronos数据格式安装依赖pip install baostock pandas使用方法1. 下载指定股票数据# 进入到data_util目录cd data_util# 下载单只股票 5分钟数据python min5_csv.py --stocks sh.600977# 下载多只股票python min5_csv.py --stocks sh.600977 sz.000001 sh.600519# 指定日期范围## 5分钟数据python min5_csv.py --stocks sh.600977 --start_date 2025-01-01 --end_date 2025-10-31## 15分钟数据python min15_csv.py --stocks sh.600977 --start_date 2025-01-01 --end_date 2025-10-31## 30分钟数据python min30_csv.py --stocks sh.600977 --start_date 2025-01-01 --end_date 2025-10-31## 60分钟数据python min60_csv.py --stocks sh.600977 --start_date 2025-01-01 --end_date 2025-10-31## 日线数据,要下载3年的数据,否则数据量不够,会报错python daily_csv.py --stocks sh.600977 --start_date 2023-01-01 --end_date 2025-10-312. 下载沪深300全部股票# 下载沪深300所有股票(默认最近7天)python min5_csv.py --hs300# 指定日期范围python min5_csv.py --hs300 --start_date 2025-10-01 --end_date 2025-10-313. 命令行参数说明参数    简写    说明    示例--stocks    -s    指定股票代码(可多个)    sh.600977--start_date    -sd    开始日期 (YYYY-MM-DD)    2025-01-01--end_date    -ed    结束日期 (YYYY-MM-DD)    2025-01-31--hs300    无    下载沪深300所有股票    -输出文件格式下载的CSV文件将保存在 ../data/ 目录下,文件命名格式为:XSHG_5min_600977.csv  # 上海股票XSHE_5min_000001.csv  # 深圳股票CSV文件格式timestamps,open,high,low,close,volume,amount2024-06-18 11:15:00,11.27,11.28,11.26,11.27,379.0,427161.02024-06-18 11:20:00,11.27,11.28,11.27,11.27,277.0,312192.0...字段说明:timestamps: 时间戳 (YYYY-MM-DD HH:MM:SS)open: 开盘价high: 最高价low: 最低价close: 收盘价volume: 成交量 (手)amount: 成交额 (千元)使用示例示例1:下载贵州茅台最近一个月数据python min5_csv.py --stocks sh.600519 --start_date 2025-08-18 --end_date 2025-09-18示例2:下载多只知名股票python min5_csv.py --stocks sh.600519 sh.600036 sz.000001 sz.000002示例3:下载沪深300全部股票的2025年1月数据python min5_csv.py --hs300 --start_date 2025-01-01 --end_date 2025-01-31注意事项股票代码格式: 必须使用 sh. 或 sz. 前缀交易时间: 只能下载交易日数据,非交易日会自动跳过数据量限制: 单次下载建议不超过3个月数据,避免数据量过大频率限制: 工具已内置防频率限制机制,请勿修改延时参数存储空间: 沪深300全部股票数据量较大,请确保有足够存储空间错误处理工具包含完善的错误处理机制:自动重试网络请求失败的情况跳过无效或格式错误的数据行处理股票代码格式错误验证日期格式有效性
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包

打赏作者

henrylin9999

你的鼓励将是我创作的最大动力

¥1 ¥2 ¥4 ¥6 ¥10 ¥20
扫码支付:¥1
获取中
扫码支付

您的余额不足,请更换扫码支付或充值

打赏作者

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值