☰
OpenChatKit 实战指南:从 Pythia-Chat-Base-7B 推理到 Llama-2-7B-32K 长上下文微调与检索增强
2026/9/25 11:51:35 网站建设 项目流程
  • 人工智能
  • 大模型
  • NLP
  • 模型训练
  • 模型推理服务

【免费下载链接】OpenChatKit

项目地址:https://gitcode.com/gh_mirrors/op/OpenChatKit
点击查看免费下载

OpenChatKit 是一个开源的对话模型工具包,仓库提供了指令微调的语言模型(Pythia-Chat-Base-7B、GPT-NeoXT-Chat-Base-20B、Llama-2-7B-32K-beta)、一个安全审核模型,以及一个可扩展的检索系统,用于从自定义知识库中获取最新上下文。本文以仓库根目录 README.md 为骨架,结合源码逐步讲解:如何用命令行工具与 Pythia-Chat-Base-7B 对话、如何复现该 7B 对话模型、如何微调长上下文模型 Llama-2-7B-32K-beta,以及如何开启检索增强(Retrieval-Augmented)能力。读完本文,你将能独立完成从环境搭建、数据/权重下载、分布式微调、权重格式转换到推理与监控的完整闭环。

OpenChatKit 中的模型基于 OIG-43M 训练数据集微调而来,该数据集由 Together、LAION 和 Ontocord.ai 三方协作构建。仓库中的代码覆盖五个核心能力:训练 20B 参数对话模型 GPT-NeoXT-Chat-Base-20B(见 docs/GPT-NeoXT-Chat-Base-20B.md)、微调 7B 参数长上下文模型 Llama-2-7B-32K-beta、训练 7B 参数对话模型 Pythia-Chat-Base-7B、用任一对话模型测试推理,以及用检索索引为模型补充额外上下文。

环境准备与快速开始

系统要求与依赖安装

官方建议使用 Miniconda + Git LFS 搭建环境。完整依赖清单在仓库根目录的 environment.yml 中,包含pytorch=2.0.1、pytorch-cuda=11.8、cudatoolkit=11.8.0、transformers=4.31.0、accelerate=0.21.0、faiss-gpu=1.7.2、wandb=0.15.5、loguru=0.6.0等。安装步骤如下:

  1. 从官网安装 Miniconda 与 Git LFS;
  2. 安装 git-lfs 钩子:
    git lfs install
  3. 在base环境安装 mamba(供所有环境复用):
    conda install mamba -n base -c conda-forge
  4. 用 mamba 创建 OpenChatKit 环境(官方注明 mamba 比 conda 快得多):
    mamba env create -f environment.yml
  5. 激活环境:
    conda activate OpenChatKit

值得注意的是,data/prepare_data.py 中的is_git_lfs_installed()会通过git lfs version检测 Git LFS 是否可用——数据集与权重下载依赖 LFS 拉取大文件,因此第 2、3 步是数据准备的前提。

与 Pythia-Chat-Base-7B 对话

Pythia-Chat-Base-7B 是 Eleuther AI 的 Pythia-6.9B-deduped 的 7B 参数指令微调版本,预训练权重以 Apache 2.0 协议发布在 Hugging Face 上。仓库提供 inference/bot.py 作为命令行测试工具,它是一个交互式 shell:输入文本,模型即回复,并自动维护对话历史作为上下文。从仓库根目录启动:

python inference/bot.py --model togethercomputer/Pythia-Chat-Base-7B

模型加载需要一些时间,加载完成后会出现欢迎提示,输入Hello开始对话:

$ python inference/bot.py Loading /home/csris/src/github.com/togethercomputer/OpenChatKit/inference/../huggingface_models/GPT-NeoXT-Chat-Base-20B to cuda:1... Welcome to OpenChatKit shell. Type /help or /? to list commands. >>> Hello. Hello human. >>>

从源码看,shell 的交互逻辑封装在OpenChatKitShell(继承自 Python 标准库cmd.Cmd):输入行若以/开头则去掉前缀作为命令执行,否则自动加上say前缀进入对话;do_say会把用户输入与模型回复依次压入Conversation(见 inference/conversation.py),再以完整历史拼出的get_raw_prompt()作为下一次生成的输入。shell 还内置了/help、/raw_say、/raw_prompt、/reset、/hyperparameters、/quit等以/开头的命令,分别用于查看帮助、无历史推理、查看完整 prompt、重置对话、查看超参数和退出(/quit由do_quit返回True结束cmdloop)。

关于命令行参数、多 GPU/指定 GPU 运行及消费级硬件部署的细节,见 inference/README.md。

推理参数详解

inference/bot.py 的main()定义了如下参数:

参数含义默认值
--gpu-id承载输入的主 GPU 设备 ID0
--model模型名称或本地路径../huggingface_models/Pythia-Chat-Base-7B
--max-tokens最大生成 token 数128
--sample是否采样生成True
--temperature采样温度0.6
--top-ktop-k 采样40
--retrieval用检索索引增强查询False
-g / --gpu-vramID:RAM形式的 GPU 显存分配(gpu-id必须出现在列表中)无
-r / --cpu-ram模型装不下 GPU 时的 CPU 内存溢出分配(GiB)无
--no-stream关闭 token 流式输出关闭流式

-g支持多个值,例如-g ID_0:RAM_0 ID_1:RAM_1 ID_N:RAM_N;max_memory字典由-g解析而来,若同时给出-r,还会把cpu键加入max_memory。推理底层由 inference/bot.py 的do_inference完成:它用StopWordsCriteria(继承transformers.StoppingCriteria)在生成的文本中出现<human>标记时提前停止,并通过stream_callback实现流式输出;生成时以pad_token_id=eos_token_id兜底,输出会去掉 prompt 前缀后再返回。

推理硬件要求与多 GPU 部署

根据 inference/README.md,GPT-NeoXT-Chat-Base-20B 至少需要 41GB 空闲显存,且每个 prompt 还会额外占用约 100–200MB:

  • 推荐至少 80GB 总显存;
  • 追求快速响应,推荐至少 48GB 单卡显存;
  • 默认仅使用 CUDA 0 号设备,推理至少需要 1 张 GPU。

多 GPU 部署时,用-g ID0:MAX_VRAM ID1:MAX_VRAM ...指定每张卡可加载的最大显存(GiB)。例如 4 张 48GB 卡可写-g 0:10 1:12 2:12 3:12 4:12,模型会先填满第一张卡再依次溢出到后续卡。需要注意:

  • MAX_VRAM仅用于加载模型,不包含后续输入,建议每张卡至少留 1–2GiB 余量、主卡(--gpu-id指定)至少留 3GiB;
  • 出现 CUDA OOM 时应调低MAX_VRAM;
  • 所有设备MAX_VRAM之和必须大于模型体积,否则bot.py会把剩余部分自动卸载到 RAM 和磁盘(此时会占满可用内存)。

指定特定 GPU 时,--gpu-id指定的设备必须同时出现在-g列表中,否则报错。示例:设备 2、5 各分配 25GiB 且以 5 为主设备,用--gpu-id 5 -g 2:25 5:25;只用设备 1 且分配 75GiB,用--gpu-id 1 -g 1:75。

消费级硬件(单卡 <48GB、多卡合计 <48GB,或遭遇 OOM)可加-r CPU_RAM限制模型占用的 RAM,例如-g 0:12 -r 20表示 CUDA 0 加载 12GiB、RAM 加载 20GiB、其余落盘到仓库根目录的offload文件夹。注意这会显著降低推理速度。这些能力在源码中对应ChatModel.__init__(见 inference/bot.py):无max_memory时用device_map="auto"直接加载;有max_memory时先以init_empty_weights()创建空权重模型,再用infer_auto_device_map生成设备映射(no_split_module_classes=["GPTNeoXLayer"]),最终通过offload_folder="offload"支持溢出落盘。

微调 Llama-2-7B-32K-beta:长上下文对话模型

Llama-2-7B-32K-beta 是 7B 参数的长上下文模型(序列长度 32768),可用多文档自然问答数据集(multi-document natural questions,mqa)和 BookSum 数据集微调。

下载并转换基础模型

在仓库根目录执行 pretrained/Llama-2-7B-32K-beta/prepare.py:

python pretrained/Llama-2-7B-32K-beta/prepare.py

权重将保存到pretrained/Llama-2-7B-32K-beta/togethercomputer_Llama-2-7B-32K-beta目录。

运行微调脚本

仓库提供了两个微调脚本:training/finetune_llama-2-7b-32k-mqa.sh与training/finetune_llama-2-7b-32k-booksum.sh。

bash training/finetune_llama-2-7b-32k-mqa.sh bash training/finetune_llama-2-7b-32k-booksum.sh

从 training/finetune_llama-2-7b-32k-mqa.sh 可以看到脚本内部实际做的是:设置GLOO_SOCKET_IFNAME、NCCL_SOCKET_IFNAME为lo,以dist_clm_train.py为核心入口,分别在--cuda-id 0..7 --rank 0..7上启动 8 个 worker(world-size 8 --pipeline-group-size 8 --data-group-size 1,即纯流水线并行)。关键训练配置包括:

  • 数据集:togethercomputer/Long-Data-Collections上的natural_questions_10_200_docs.jsonl.zst:1(BookSum 脚本对应booksum.jsonl.zst:1),冒号后为采样权重;
  • --model-type llama、--num-layers 4(每 GPU 4 层)、--embedding-dim 4096;
  • --seq-length 32768、--batch-size 4、--micro-batch-size 1、--gradient-accumulate-step 1;
  • --lr 2e-5、--total-steps 10、--checkpoint-steps 10、--warmup-steps 0;
  • --fp16混合精度、--optimizer adam、--seed 42、--load-pretrained-model true、--dist-url tcp://127.0.0.1:7033;
  • --dp-backend nccl --dp-mode allreduce --pp-mode gpipe --profiling no-profiling。

三个环境变量可覆盖脚本默认值:FINETUNE_TOTAL_STEPS、FINETUNE_CHECKPOINT_STEPS、FINETUNE_CHECKPOINT_PATH(默认model_ckpts/llama-2-7b-32k-mqa)。训练过程中,checkpoint 会保存到仓库根目录的model_ckpts目录。

更多自定义训练的方法见 training/README.md。

训练参数说明(来自 training/README.md)

需要设置的环境变量:

export GLOO_SOCKET_IFNAME=lo # 需与 --net-interface 一致 export NCCL_SOCKET_IFNAME=lo # 需与 --net-interface 一致 export WANDB_NAME=gptj-test # wandb 运行名

需要仔细设置的核心参数:

  • --model-name:按层分片的模型 checkpoint 路径;
  • --tokenizer-name:通常与--model-name相同,也可用 HF 模型名;
  • --model-type:模型类型,如gptj/gptneox/llama;
  • --num-layers:每张 GPU的 Transformer 层数。例如 GPT-J 共 28 层,两台 GPU 组流水线时传 14;
  • --embedding-dim:模型隐藏层维度(GPT-J-6B 为 4096),用于创建缓冲区;
  • --dist-url:rank 0 worker 的地址,所有 worker 必须可访问。单机多卡可写--dist-url tcp://127.0.0.1:7033;
  • --world-size:worker 总数,满足world-size == pipeline-group-size *>mkdir huggingface_models \ && python tools/convert_to_hf_llama.py \ --config-name togethercomputer/Llama-2-7B-32K-beta \ --ckpt-path model_ckpts/llama-2-7b-32k-mqa/checkpoint_10 \ --save-path huggingface_models/llama-2-7b-32k-mqa \ --n-stages 4 \ --n-layer-per-stage 8 \ --fp16

    其中--fp16以 fp16 加载并存储模型;--n-stages与--n-layer-per-stage须与训练时的流水线配置匹配。请把model_ckpts/llama-2-7b-32k-mqa/checkpoint_10替换为model_ckpts/llama-2-7b-32k-mqa或model_ckpts/llama-2-7b-32k-booksum目录下的最新 checkpoint。

    从源码看(tools/convert_to_hf_llama.py),load_decentralized_checkpoint会遍历每个流水线 stage,从prank_{i}_checkpoint.pt加载权重:stage 0 处理 embedding 与首段 Transformer 层,最后一个 stage 额外加载norm.weight与lm_head.weight(训练时 embedding 与 lm_head 是分开训练的),其余 stage 加载各自层段;同时用no_init_weights关闭随机初始化以加速空模型创建。转换完成后,save_pretrained会把模型、config 与 tokenizer 一并写入--save-path。

    复现 Pythia-Chat-Base-7B

    本教程演示如何用 OIG 数据集微调 Eleuther AI 的 Pythia-6.9B-deduped,从而复现 Pythia-Chat-Base-7B。

    下载训练数据与基础模型

    OIG 数据集由 LAION、Together 与 Ontocord.ai 共建。从仓库根目录执行:

    python data/OIG/prepare.py

    数据将落在data/OIG/files目录。实现上,data/OIG/prepare.py 会调用 data/prepare_data.py 的prepare_data,以https://huggingface.co/datasets/laion/OIG为数据源克隆仓库,并对*.gz文件自动解压。prepare_data还支持本地文件、GitHub/HF 仓库、S3 兼容存储(通过-a/-k传 AWS 密钥,或从环境变量AWS_ACCESS_KEY_ID等读取)和普通 HTTP(S) URL 四类数据源。

    接着下载基础模型:

    python pretrained/Pythia-6.9B-deduped/prepare.py

    权重将保存在pretrained/Pythia-6.9B-deduped/EleutherAI_pythia-6.9b-deduped目录。该脚本(见 pretrained/Pythia-6.9B-deduped/prepare.py)复用 pretrained/prepare_pretrained.py 的prepare_pretrained:从 Hugging Face 拉取 config、tokenizer 与模型,再把 embedding(embed_in.weight)、每个 Transformer 层(pytorch_{i}.pt)和embed_out.weight+final_layer_norm(pytorch_lm_head.pt)拆分成独立文件,供训练框架按层分片加载;该文件也支持--offload-dir参数将模型卸载到磁盘以节省内存。

    (可选)8bit Adam

    如需使用 8bit-adam 优化器,先安装:

    pip install bitsandbytes # optional, to use 8bit-adam

    然后在训练参数中把--optimizer改为8bit-adam。

    训练模型

    training/finetune_Pythia-Chat-Base-7B.sh 配置并运行训练循环。下载完数据与基础模型后执行:

    bash training/finetune_Pythia-Chat-Base-7B.sh

    该脚本使用--model-type gptneox,把data/OIG/files下的 25 个任务以不同采样权重组合进--task-name,例如unified_ni.jsonl:0.2、unified_p3.jsonl:0.5、unified_flan.jsonl:0.2、unified_chip2.jsonl:0.01、unified_rallio_safety_and_prosocial.jsonl:0.1等,涵盖指令、对话、安全与亲社会、代码、摘要、问答、数学推理、创作等多种能力;训练配置为--total-steps 20000 --checkpoint-steps 100 --lr 1e-5 --seq-length 2048 --batch-size 32 --world-size 8 --pipeline-group-size 4 --data-group-size 2(2 条数据并行流水线)。训练循环把 checkpoint 保存到model_ckpts目录。

    转换权重为 Hugging Face 格式

    使用 tools/convert_to_hf_gptneox.py:

    mkdir huggingface_models \ && python tools/convert_to_hf_gptneox.py \ --config-name EleutherAI/pythia-6.9b-deduped \ --ckpt-path model_ckpts/Pythia-Chat-Base-7B/checkpoint_100 \ --save-path huggingface_models/Pythia-Chat-Base-7B \ --n-stages 4 \ --n-layer-per-stage 8 \ --fp16

    --fp16以 fp16 加载并存储模型。请把model_ckpts/Pythia-Chat-Base-7B/checkpoint_100替换为model_ckpts/Pythia-Chat-Base-7B目录下的最新 checkpoint。与 Llama 转换脚本对称,load_decentralized_checkpoint(tools/convert_to_hf_gptneox.py)逐 stage 读取prank_{i}_checkpoint.pt,首 stage 载入embed_in.weight,末 stage 载入final_layer_norm与embed_out。

    测试新模型

    用 OpenChatKit Shell 与新模型对话,默认加载huggingface_models目录下名为 Pythia-Chat-Base-7B 的模型,也可用--model覆盖:

    python inference/bot.py python inference/bot.py --model ./huggingface_models/GPT-NeoXT-Chat-Base-20B

    模型加载完成后在提示符输入文本即可对话(交互界面与上文一致),shell 同样支持查看超参数、完整 prompt 等/命令;/quit退出。

    训练监控

    默认训练脚本只打印 loss,但可以通过--train-log-backend切换日志后端。

    Loguru

    为训练脚本追加--train-log-backend loguru,指标将写入./logs/file_{time}.log。仓库在 environment.yml 中固定了loguru==0.6.0,相关实现位于 training/utils/logging_utils.py。

    Weights & Biases

    先登录:

    wandb login

    再在训练脚本中设置--train-log-backend wandb即可上报到 Weights & Biases。仓库固定wandb==0.15.5,运行名可通过环境变量WANDB_NAME指定(见 training/README.md)。训练入口 training/dist_clm_train.py 会读取该后端参数并实例化对应的报告器。

    实验性功能:检索增强模型(Retrieval-Augmented)

    警告:检索支持目前是实验性的。

    retrieval 目录实现了一个用于查询 Wikipedia 的 Faiss 索引的 Python 包。核心类WikipediaIndex在 retrieval/wikipedia.py 中:它用facebook/contriever-msmarco编码查询,通过mean_pooling得到句向量,在 Faiss 索引(knn.index,以 MMAP 只读方式加载)中检索,再对命中的句子的前后各w=5条相邻句子按余弦相似度阈值w_th=0.5扩展上下文。下面是让 bot 使用该索引增强查询的步骤。

    1. 下载 Wikipedia 索引:

      python data/wikipedia-3sentence-level-retrieval-index/prepare.py
    2. 以--retrieval标志启动 bot:

      python inference/bot.py --retrieval

    启动后 bot 会同时加载对话模型与检索索引(耗时较长),加载完成后所有查询都会被附加额外上下文:

    $ python inference/bot.py --retrieval Loading /OpenChatKit/inference/../huggingface_models/GPT-NeoXT-Chat-Base-20B to cuda:0... Loading retrieval index... Welcome to OpenChatKit shell. Type /help or /? to list commands. >>> Where is Zurich? Where is Zurich? Zurich is located in Switzerland. >>>

    这一行为在 inference/bot.py 的do_say中实现:若开启--retrieval,先用self._index.search(arg)检索,把命中文本通过push_context_turn注入对话上下文,再与用户输入拼接生成回复——这正是"用检索索引为模型补充最新上下文"的最小实现。

    更多资源与其他模型

    OpenChatKit 还提供更大的 20B 参数对话模型 GPT-NeoXT-Chat-Base-20B(基于 Eleuther AI 的 GPT-NeoXT-Chat-Base-20B 训练),训练与复现细节见 docs/GPT-NeoXT-Chat-Base-20B.md;仓库还附有 docs/finetuning-RedPajama-3B.md 等文档。微调脚本方面,training/finetune_GPT-NeoXT-Chat-Base-20B.sh、training/finetune_RedPajama-INCITE-Chat-3B-v1.sh 等可参考;LoRA 示例见 training/lora/example。通用自定义训练的完整参数说明见 training/README.md。

    许可与引用

    仓库代码由 Together Computer 开发(除非另有注明),版权归 Together Computer 所有,以 Apache 2.0 协议开源(完整条款见 LICENSE)。若在论文或产品中引用 OpenChatKit,可参考以下 BibTeX:

    @software{openchatkit, title = {{OpenChatKit: An Open Toolkit and Base Model for Dialogue-style Applications}}, author = {Together Computer}, url = {https://github.com/togethercomputer/OpenChatKit} month = {3}, year = {2023}, version = {0.15}, }

    结语

    本文从仓库根目录 README 出发,完整走通了 OpenChatKit 的两条主线:一是"下载即用"的推理路径(Pythia-Chat-Base-7B + bot.py + 可选的 Faiss 检索增强),二是"从零复现/微调"的训练路径(数据与权重准备 → 分布式微调 → checkpoint 转 HF 格式 → 推理验证)。配合仓库中的微调脚本、转换工具与源码注释,你可以在此基础上替换自己的jsonl数据集、调整采样权重与并行规模,微调出面向特定领域的对话模型。需要注意的是:长上下文(32K)微调与 20B 模型推理对 GPU 显存和内存有较高要求,消费级硬件请务必配合-g/-r参数规划好资源分配。

    • 人工智能
    • 大模型
    • NLP
    • 模型训练
    • 模型推理服务

    【免费下载链接】OpenChatKit

    项目地址:https://gitcode.com/gh_mirrors/op/OpenChatKit
    点击查看免费下载

    相关推荐

    上一篇:蘑菇书实战:使用 Double-DQN 从零实现 CartPole-v0 平衡控制
    下一篇:京东抢购助手终极教程:3步实现Python自动化下单

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询