首页
/ Qlib Task Management 详解:从任务生成、MongoDB 存储到多进程训练与结果集成的完整流程

Qlib Task Management 详解:从任务生成、MongoDB 存储到多进程训练与结果集成的完整流程

2026-09-05 12:26:29作者:魏献源Searcher

本文基于 Qlib 仓库文档 docs/advanced/task_management.rst 展开。Qlib 的 qrun 一次只能执行单个任务,而真实的量化研究中往往需要"同一模板 × 多个时间段 × 多个模型/损失函数"地批量实验。Task Management 模块正是为此设计的:它提供 任务生成(TaskGen)→ 任务存储(TaskManager/MongoDB)→ 任务训练(Trainer/run_task)→ 结果集成(Collector/Group/Ensemble) 的全链路能力,并可直接服务于在线更新(Online Serving)等滚动训练场景。读完本文,你将能够:理解 RollingGen 等任务生成器的参数与滚动语义、配置 TaskManager 并使用其状态机管理任务生命周期、用 TrainerR/TrainerRM 完成(并行)训练,以及用 RecorderCollector + RollingGroup + RollingEnsemble 把分散的 Recorder 结果拼回完整的时间序列。

Qlib 任务生成、训练与结果收集的整体流程图

一、整体架构:四个环节与它们之间的数据流

整个流程可概括为一条流水线:

  1. Task Generating:基于固定的 task 模板,用可自定义的 TaskGen 派生出大量任务(不同时间段、不同损失、不同模型);
  2. Task Storing:把所有任务写入 MongoDB,由 TaskManager 以状态机方式管理任务生命周期,支持并发、错误处理与集群化执行;
  3. Task Training:从任务池中取出 WAITING 状态的任务并执行训练,默认执行器是 qlib.model.trainer.task_train,会完整跑完 task 中定义的 Model、Dataset、Record;
  4. Task Collecting:训练完成后,用 CollectorGroupEnsemble 以"可读、可扩展、松耦合"的方式收集并集成各 Recorder 的产物。

一个可运行的端到端示例位于 examples/model_rolling/task_manager_rolling.py,后文第五节会逐段拆解它。

二、Task Generating:用一个模板批量派生任务

2.1 任务模板与 TaskGen 基类

一个 Qlib taskModelDatasetRecord(以及用户自行扩展的任意字段)组成,其标准结构参见工作流文档 docs/component/workflow.rst 中的 Task Section。TaskGen 的基类定义在 qlib/workflow/task/gen.py,只有一个抽象方法:

class TaskGen(metaclass=abc.ABCMeta):
    @abc.abstractmethod
    def generate(self, task: dict) -> List[dict]:
        """基于一个 task 模板生成若干任务"""

    def __call__(self, *args, **kwargs):
        return self.generate(*args, **kwargs)

也就是说,自定义生成器只需实现"输入一个任务模板 dict、输出一批任务 dict"的映射。task_generator 函数(qlib/workflow/task/gen.py#L16-L50)进一步支持多模板 × 多生成器的笛卡尔积

tasks = task_generator(
    tasks=[task_a, task_b, task_c],   # N 个模板
    generators=[gen_A, gen_B],        # 每个生成器对每个模板产生若干变体
)
# 若 gen_A 对每个模板产生 2 个、gen_B 产生 3 个,则最终生成 3 * 2 * 3 = 18 个任务

2.2 RollingGen:按滚动窗口派生任务

RollingGenqlib/workflow/task/gen.py#L140-L301)是最核心的生成器,它把一个任务模板按时间轴滚动展开,让用户"在单个实验中验证不同时期数据对模型的影响"。构造函数参数如下(均取自源码签名与 docstring):

参数 默认值 说明
step 40 滚动步长(交易日数),每个后续任务的 test 段向后平移 step
rtype ROLL_EX(expanding) 滚动类型:ROLL_EX 为"起始固定、结束扩展",ROLL_SD 为"窗口大小固定、整体滑动"(对应 TimeAdjuster.SHIFT_EX / SHIFT_SD
ds_extra_mod_func handler_mod 生成每个任务后的额外修改钩子,默认为 handler_mod
test_key / train_key "test" / "train" segments 中标记测试段/训练段的 key
trunc_days None 截断若干天以避免未来信息泄漏(见 2.3)
task_copy_func copy.deepcopy 深拷贝整个任务;若希望任务间共享某些对象可自定义

其滚动语义(generate + gen_following_tasksqlib/workflow/task/gen.py#L187-L301)值得注意两点:

  • 首个滚动任务:test 段被重设为以原 test 段起点开始、长度恰为 step 的窗口(test_start ~ test_start + step - 1 个交易日),即每个任务对应一个等长的测试期;
  • 后续滚动:逐段调用 TimeAdjuster.shift 平移。在 ROLL_EX 模式下只有 train 段扩展(起点不动、终点后移),valid 与 test 段保持固定尺寸滑动;当 test 段起点越过原始 test 终点时停止。

TimeAdjusterqlib/workflow/task/utils.py#L82-L280)负责所有日期对齐工作:align_seg 把任意日期对齐到交易日历,shift 按"交易日索引"平移而非自然日,保证跨节假日的滚动也严格等长。

2.3 两个防泄漏细节:handler_mod 与 trunc_segments

  • handler_modqlib/workflow/task/gen.py#L94-L123):滚动后 handler 的 end_time 可能早于新 test 段的终点,导致 handler 读不到足够的数据。该钩子会在 test 段终点是 None(开放式"至今")或早于 handler end_time 时,把 handler 的 end_time 同步扩展为 test 段终点。
  • trunc_segmentsqlib/workflow/task/gen.py#L126-L137):当设置了 trunc_days 时,train/valid 段的终点会被强制截断到 test_start - trunc_days,防止训练集吃到 test 期标签所依赖的"近未来"数据。

此外,仓库还实现了 MultiHorizonGenBaseqlib/workflow/task/gen.py#L304-L350):给定 horizon 列表(预测周期)与 label_leak_n(预测日之后需等待多少个未来交易日标签才完整,例如 Ref($close, -2)/Ref($close, -1) - 1 这类标签取 2),为每个 horizon 生成一份任务,并自动按 horizon + label_leak_n 截断 segments——这是把"多周期预测实验"也纳入同一套 TaskGen 框架的示例。

三、Task Storing:TaskManager 与 MongoDB

3.1 前置条件:必须先配置 MongoDB

文档明确强调:使用 TaskManager 之前,用户必须先完成 MongoDB 配置。有两种等价方式(参见 docs/start/initialization.rstqlib.initmongo 参数说明):

# 方式一:在 qlib.init 时传入
qlib.init(provider_uri=provider_uri, region=region, mongo={
    "task_url": "mongodb://localhost:27017/",   # your MongoDB url
    "task_db_name": "rolling_db",               # database name
})

# 方式二:init 之后直接写全局配置
from qlib.config import C
C["mongo"] = {
    "task_url": "mongodb://localhost:27017/",
    "task_db_name": "rolling_db",
}

读取逻辑在 qlib/workflow/task/utils.py#L22-L57get_mongodb:未配置 C["mongo"] 会直接报错;配置中的 task_url 支持带凭据的形式(如 mongodb://user:pwd@host:port,见 docs/start/initialization.rst)。

3.2 任务文档结构与四种状态

TaskManager 定义在 qlib/workflow/task/manage.py。每个 TaskManager(task_pool) 实例对应 MongoDB 中的一个 Collection(task_pool 即集合名),任务文档结构(源码 docstring)为:

{
    "def": pickle 序列化后的任务定义(任务体本身),
    "filter": json-like 字段,用于按内容检索任务(去重依据),
    "status": "waiting" | "running" | "part_done" | "done",
    "res": pickle 序列化后的任务结果,
}

四个状态常量(qlib/workflow/task/manage.py#L79-L82):

  • STATUS_WAITING:等待训练;
  • STATUS_RUNNING:正在训练;
  • STATUS_PART_DONE:完成了部分步骤、等待下一步(DelayTrainer 两步式训练的关键);
  • STATUS_DONE:全部完成。

def/res 字段用 pickle + Binary 编码进出数据库(_encode_task/_decode_task),读取时使用 restricted_pickle_loads 做受限反序列化,这是一个值得注意的安全细节。

3.3 核心 API 一览(按生命周期排序)

以下方法签名与语义均来自 qlib/workflow/task/manage.py 源码:

方法 作用
create_task(task_def_l, dry_run, print_nt) 批量入库:按 filter 查重,新任务以 waiting 插入,已存在任务只返回其 _id;支持 dry_run 只统计不写入
fetch_task(query, status) 原子地取出一个指定状态任务并置为 running,内部用 find_one_and_updatepriority 降序排序,保证多进程/多机器并发时每个任务只被取走一次
safe_fetch_task 上下文管理器版取任务:执行期间抛异常或收到 KeyboardInterrupt 时自动把任务状态回滚为原状态(return_task
task_fetcher_iter(query) 循环 safe_fetch_task 的迭代器,任务池取空后结束
commit_task_res(task, res, status) 写回结果到 res 字段并设置终态(默认 done,也可设为 part_done
return_task(task, status) 错误处理用,把任务状态退回(默认 waiting
wait(query) 阻塞轮询(每 10 秒统计一次),直到所有任务不再是 waiting/running/part_done;多进程场景下主进程用它等待其他进程/机器完成
task_stat(query) / re_query(_id) / reset_waiting / reset_status / prioritize(task, priority) / remove(query) 状态统计、按 _id 回查(返回 res 即 Recorder)、把意外退出的 running 重置回 waiting、设置优先级、删除任务

由于 fetch_task 基于 MongoDB 的原子更新实现,"取一个任务并标记 running"天然互斥——这正是文档所说的"自动抓取未完成任务、带错误处理地管理一组任务生命周期",也是集群化(多机抢任务)能成立的基础。

TaskManager 本身还是 CLI(模块末尾 fire.Fire(TaskManager)),常用命令:

python -m qlib.workflow.task.manage -h                 # 查看 manage 模块 CLI 手册
python -m qlib.workflow.task.manage wait -h            # 查看 wait 子命令手册
python -m qlib.workflow.task.manage -t <pool_name> wait        # 等待任务池清空
python -m qlib.workflow.task.manage -t <pool_name> task_stat  # 查看各状态任务数

四、Task Training:run_task 与四种 Trainer

4.1 run_task:任务池的消费循环

run_taskqlib/workflow/task/manage.py#L485-L551)是"从池中不断取任务并执行"的通用驱动:

run_task(
    task_func,      # def (task_def, **kwargs) -> res;最简取法就是 qlib.model.trainer.task_train
    task_pool,      # 任务池(MongoDB Collection)名
    query={},       # 只消费符合 query 的任务
    force_release=False,   # True 时用单进程 ProcessPoolExecutor 执行,便于强制释放内存
    before_status=TaskManager.STATUS_WAITING,  # 可传 STATUS_PART_DONE 以续跑两步式任务
    after_status=TaskManager.STATUS_DONE,
    **kwargs,       # 透传给 task_func(如 experiment_name)
)

其状态迁移规则(源码 docstring):

  • WAITING -> DONEWAITING -> PART_DONE:取 task["def"] 作为 task_func 入参;
  • PART_DONE -> PART_DONEPART_DONE -> DONE:取 task["res"] 作为入参(中间结果已存入 res)。

默认的 task_func 使用 qlib.model.trainer.task_trainqlib/model/trainer.py#L108-L128)即可,它在一个 Recorder 中完成整个工作流:记录任务配置 → 按配置实例化并 fit Model → 保存 params.pkldatasetdump_all=False 以便在线推理复用)→ 依次生成 task 中声明的各 Record(预测、回测、分析)。

4.2 Trainer 家族:TrainerR 与 TrainerRM

Trainer 基类(qlib/model/trainer.py#L131-L206)定义 train(训练任务列表并返回 Recorder 列表)与 end_train(收尾)两步,__call__ 等价于 end_train(train(...))。仓库提供两种主要实现,以及各自的延迟(Delay)变体:

  • TrainerRqlib/model/trainer.py#L209-L290):最简单的方式,线性地逐个训练任务列表,默认 train_func=task_train,可选 call_in_subproc=True 在子进程中执行以强制释放内存。不需要 MongoDB——如果不想用 TaskManager 管理生命周期,用它训练 TaskGen 生成的任务列表就足够(这是文档给出的选择建议)。
  • TrainerRMqlib/model/trainer.py#L341-L488):基于 TaskManager 的"记录器 + 任务池"训练器。traincreate_task 把全部任务写入 MongoDB,再调 run_task 只消费本批任务(query={"_id": {"$in": _id_list}}),随后 tm.wait 等待全部完成,最后按 _id 回查 res 得到 Recorder 列表(并打上 train_status_id in TaskManager 标签)。它还实现了 worker():与 train 共享同一 task_pool、可运行在其他进程甚至其他机器上,从而把"提交"与"执行"解耦。
  • DelayTrainerR / DelayTrainerRMqlib/model/trainer.py#L293-L338qlib/model/trainer.py#L491-L619):train 阶段只执行 begin_task_train(创建 Recorder、保存任务配置),把耗时的真实 fit 推迟到 end_train 阶段的 end_task_train 执行;对应到任务池上,任务先流转到 PART_DONEend_train 再以 before_status=PART_DONE 消费。文档指出这个"DelayTrainer"概念可用于在线模拟中的并行训练;DelayTrainerRMskip_run_task=True 则支持"在 CPU 机器上提交任务、等待 GPU 机器上的 worker 完成"这类异构部署。

五、Task Collecting:Collector / Group / Ensemble 三级结构

收集结果前,文档要求先通过 qlib.init 指定 mlruns 路径(结果都存放在 MLflow Recorder 中)。三个角色及其层级关系(文档原文)是:

  • Collector:两步——collect(把任何对象收集成 dict)与 process_collect(按 process_list 依次处理该 dict);
  • Group:两步——group(按 group_func 把对象集合分组为 dict)与 reduce(按规则把每个分组的 dict 归约为单一对象)。例如 {(A,B,C1): object, (A,B,C2): object} --group--> {(A,B): {C1: object, C2: object}} --reduce--> {(A,B): object}
  • Ensemble:把组内对象合并为一个对象,{C1: object, C2: object} --Ensemble--> object

层级对应关系:Collector 的第二步(process_collect)对应 Group,Group 的第二步(reduce)对应 Ensemble

源码印证(qlib/workflow/task/collect.py#L19-L87qlib/model/ens/group.pyqlib/model/ens/ensemble.py):

  • Collector.process_collect 对收集到的每个 artifact 依次执行 process_list 中的可调用对象;
  • Group.__call__group 分组,再用 joblib.Parallel(n_jobs=...) 对各组并行 reduce(即调用内部 Ensemble),因此 Group(group_func, ens=Ensemble) 就是文档例子的完整实现;
  • 常用 Ensemble 有:AverageEnsembleqlib/model/ens/ensemble.py#L91-L132,对同一时间段不同模型的预测/IC 做逐日标准化后取均值)、RollingEnsembleqlib/model/ens/ensemble.py#L65-L88,把相邻滚动窗口的结果按 datetime 索引拼接成完整时间序列,重复处保留最新预测)以及用于结果扁平化展示的 SingleKeyEnsemble
  • 最常用的具体收集器是 RecorderCollectorqlib/workflow/task/collect.py#L136-L253):从某 Experiment 的所有 Recorder 中按 rec_key_func 生成键、按 rec_filter_func 过滤,再按 artifacts_path(默认 {"pred": "pred.pkl"})加载产物,输出 {artifact: {rec_key: object}}

六、端到端示例:examples/model_rolling/task_manager_rolling.py

examples/model_rolling/task_manager_rolling.py 是文档推荐的"全流程示例",用 RollingTaskExample 类(fire 驱动)串起四个环节,值得逐段对照阅读:

class RollingTaskExample:
    def __init__(self,
        provider_uri="~/.qlib/qlib_data/cn_data",
        region=REG_CN,
        task_url="mongodb://10.0.0.4:27017/",
        task_db_name="rolling_db",
        experiment_name="rolling_exp",
        task_pool=None,          # 指定后走 TrainerRM(MongoDB),否则走 TrainerR
        task_config=None,        # 默认用 [XGBOOST_TASK_CONFIG, LGB_TASK_CONFIG] 两个模板
        rolling_step=550,
        rolling_type=RollingGen.ROLL_SD,
    ):
        mongo_conf = {"task_url": task_url, "task_db_name": task_db_name}
        qlib.init(provider_uri=provider_uri, region=region, mongo=mongo_conf)
        # task_pool 为 None -> TrainerR;否则 TrainerRM(experiment_name, task_pool)
        self.rolling_gen = RollingGen(step=rolling_step, rtype=rolling_type)

四个环节对应四个方法:

def task_generating(self):
    # 2 个任务模板 × RollingGen(550 步滑动) -> 若干滚动任务
    return task_generator(tasks=self.task_config, generators=self.rolling_gen)

def task_training(self, tasks):
    self.trainer.train(tasks)          # TrainerR 或 TrainerRM 的 train

def worker(self):
    # 仅 TrainerRM:在其他进程/机器上消费任务池,等价于 TrainerRM.worker
    run_task(task_train, self.task_pool, experiment_name=self.experiment_name)

def task_collecting(self):
    def rec_key(recorder):             # 用 (模型类名, test 段) 作为收集键
        task_config = recorder.load_object("task")
        return task_config["model"]["class"], \
               task_config["dataset"]["kwargs"]["segments"]["test"]

    def my_filter(recorder):          # 只保留 LGBModel 的结果
        return rec_key(recorder)[0] == "LGBModel"

    collector = RecorderCollector(
        experiment=self.experiment_name,
        process_list=RollingGroup(),  # 按 (模型, 滚动段) 分组 + RollingEnsemble 拼接
        rec_key_func=rec_key,
        rec_filter_func=my_filter,
    )
    print(collector())

main 依次执行 reset -> task_generating -> task_training -> task_collecting;命令行入口为 fire.Fire(RollingTaskExample),例如 python task_manager_rolling.py main --experiment_name="your_exp_name"。注意其中的两处设计:reset 会清空任务池与实验内所有 Recorder(示例脚本的破坏性操作,正式环境需谨慎);RollingGroup 默认 ens=RollingEnsemble(),依赖"滚动键位于 key 元组末尾"的约定把键为 (模型, test段) 的结果先按模型分组、再拼接成跨窗口的完整预测序列。

七、实践要点小结

  1. 先选"要不要 MongoDB":只做一次性批量实验时用 TaskGen + TrainerR 即可,零外部依赖;需要任务去重、失败重试、多机并发或在线滚动更新时,再引入 TaskManager + TrainerRM,且必须先配置 C["mongo"]task_url + task_db_name)。
  2. 滚动参数决定实验语义step 控制测试窗口长度与滚动频率,rtype 控制是扩展训练集(ROLL_EX)还是固定窗口滑动(ROLL_SD);跨节假日的等长保证由 TimeAdjuster 的交易日索引平移提供。
  3. 防泄漏是滚动任务的默认行为handler_mod 同步 handler 数据终点、trunc_days/MultiHorizonGenBaselabel_leak_n 截断训练段,避免"测试期近未来信息"渗入训练/验证集。
  4. PART_DONE 是并行化与延迟训练的关键DelayTrainer*WAITING -> PART_DONE -> DONE 两段式把"任务准备"与"真实 fit"解耦,配合 worker() 可实现提交与执行分离的异构部署;fetch_task 的原子更新保证多机抢任务不重复。
  5. 收集端保持键的稳定性RecorderCollectorrec_key_func 需要能唯一且可分组地描述每个 Recorder(示例用"模型类名 + test 段"),Group/Ensemble 才能正确地把滚动结果按模型归组并按时间拼接。

文中涉及的实现与文档路径索引:docs/advanced/task_management.rstqlib/workflow/task/gen.pyqlib/workflow/task/manage.pyqlib/workflow/task/collect.pyqlib/workflow/task/utils.pyqlib/model/trainer.pyqlib/model/ens/group.pyqlib/model/ens/ensemble.pyexamples/model_rolling/task_manager_rolling.pydocs/start/initialization.rst

登录后查看全文
热门项目推荐
相关项目推荐