初始化项目,由ModelHub XC社区提供模型
Model: XingChina/ChunMengDie-1.0-0.4b Source: Original Platform
This commit is contained in:
38
.gitattributes
vendored
Normal file
38
.gitattributes
vendored
Normal file
@@ -0,0 +1,38 @@
|
|||||||
|
*.7z filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.arrow filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.bin filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ftz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.gz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.h5 filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.joblib filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.model filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.npy filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.npz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.onnx filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ot filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.parquet filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pb filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pickle filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pkl filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pt filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pth filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.rar filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
||||||
|
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tar filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tflite filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tgz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.wasm filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.xz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.zip filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.zst filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
chunmengdie_f16\[1\].gguf filter=lfs diff=lfs merge=lfs -text
|
||||||
|
chunmengdie_q4_km\[1\].gguf filter=lfs diff=lfs merge=lfs -text
|
||||||
|
chunmengdie_q8_0\[1\].gguf filter=lfs diff=lfs merge=lfs -text
|
||||||
28
LICENSE
Normal file
28
LICENSE
Normal file
@@ -0,0 +1,28 @@
|
|||||||
|
Copyright (c) 2026, XingChina
|
||||||
|
All rights reserved.
|
||||||
|
|
||||||
|
Redistribution and use in source and binary forms, with or without
|
||||||
|
modification, are permitted provided that the following conditions are met:
|
||||||
|
|
||||||
|
1. Redistributions of source code must retain the above copyright notice,
|
||||||
|
this list of conditions and the following disclaimer.
|
||||||
|
|
||||||
|
2. Redistributions in binary form must reproduce the above copyright notice,
|
||||||
|
this list of conditions and the following disclaimer in the documentation
|
||||||
|
and/or other materials provided with the distribution.
|
||||||
|
|
||||||
|
3. Neither the name of the copyright holder nor the names of its
|
||||||
|
contributors may be used to endorse or promote products derived from
|
||||||
|
this software without specific prior written permission.
|
||||||
|
|
||||||
|
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||||
|
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||||
|
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
|
||||||
|
ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE
|
||||||
|
LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
|
||||||
|
CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
|
||||||
|
SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
|
||||||
|
INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
|
||||||
|
CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
|
||||||
|
ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
|
||||||
|
POSSIBILITY OF SUCH DAMAGE.
|
||||||
189
README.md
Normal file
189
README.md
Normal file
@@ -0,0 +1,189 @@
|
|||||||
|
---
|
||||||
|
license: bsd-3-clause
|
||||||
|
datasets:
|
||||||
|
- BelleGroup/train_0.5M_CN
|
||||||
|
- liumindmind/NekoQA-10K
|
||||||
|
- cyberlangke/Nana-catgirl-dataset-110k
|
||||||
|
- XingChina/ChunMengDie-1.0-User-Data
|
||||||
|
language:
|
||||||
|
- zh
|
||||||
|
- en
|
||||||
|
pipeline_tag: text-generation
|
||||||
|
library_name: transformers
|
||||||
|
tags:
|
||||||
|
- chinese
|
||||||
|
- instruction-tuning
|
||||||
|
- conversational
|
||||||
|
- experimental
|
||||||
|
- safetensors
|
||||||
|
- gguf
|
||||||
|
- quantized
|
||||||
|
- 8bit
|
||||||
|
- 4bit
|
||||||
|
- llama.cpp
|
||||||
|
---
|
||||||
|
|
||||||
|
# ChunMengDie-1.0-0.4B
|
||||||
|
|
||||||
|
## 关于作者
|
||||||
|
|
||||||
|
一个热爱 AI 的八年级学生。欢迎交流学习。
|
||||||
|
|
||||||
|
## 模型简介
|
||||||
|
|
||||||
|
ChunMengDie-1.0-0.4B 是一个基于 Transformer 架构的中文对话实验模型,参数量 4 亿(0.4B)。
|
||||||
|
|
||||||
|
**当前版本状态**:
|
||||||
|
- 早期研究阶段,输出质量不稳定,存在大量乱码和语义不连贯现象。
|
||||||
|
- 仅用于学术探索和实验验证,**不建议用于任何生产环境**。
|
||||||
|
- 后续版本将持续迭代优化。
|
||||||
|
- 由于其他分词器死活下不下来,使用内置的gpt2分词器进行训练(我自己尝试过训练自定义分词器,但是llama.cpp转换出现错误,尝试把分词器加进llama.cpp也不行,所以选择内置的分词器)
|
||||||
|
- 使用动态学习率
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 📂 仓库文件结构
|
||||||
|
|
||||||
|
本仓库同时提供 **Hugging Face 原生格式** 和 **llama.cpp GGUF 量化格式** 两种权重,请根据你的部署场景选择:
|
||||||
|
|
||||||
|
### 🤗 Hugging Face 格式(适用于 `transformers` 库)
|
||||||
|
|
||||||
|
| 文件名 | 说明 |
|
||||||
|
|--------|------|
|
||||||
|
| `config.json` | 模型结构配置文件 |
|
||||||
|
| `generation_config.json` | 文本生成参数配置 |
|
||||||
|
| `model.safetensors` | BF16 原始精度权重(约 0.8 GB) |
|
||||||
|
| `tokenizer.json` | 分词器词表 |
|
||||||
|
| `tokenizer_config.json` | 分词器配置 |
|
||||||
|
|
||||||
|
### 🦙 llama.cpp GGUF 格式(适用于 `llama.cpp` 或 `ollama` 部署)
|
||||||
|
|
||||||
|
| 精度 | 文件名 | 文件大小 | 适用场景 |
|
||||||
|
|------|--------|----------|----------|
|
||||||
|
| **BF16(原始)** | `chunmengdie_f16.gguf` | 501 MB | 追求最佳效果,适合高端显卡 |
|
||||||
|
| **8bit 量化** | `chunmengdie_q8_0.gguf` | 275 MB | 常规推理,质量损失极小(推荐) |
|
||||||
|
| **4bit 量化** | `chunmengdie_q4_km.gguf` | 181 MB | 低显存设备(如 4GB 显卡)或移动端部署 |
|
||||||
|
|
||||||
|
> 💡 **选择建议**:如果你是 Python 开发者,用 `transformers` 加载 HF 格式最方便;如果你在本地命令行或边缘设备部署,用 GGUF 格式配合 `llama.cpp` 更高效。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 🚀 推理使用指南
|
||||||
|
|
||||||
|
### 方式一:Transformers(HF 格式)
|
||||||
|
|
||||||
|
安装依赖:
|
||||||
|
```bash
|
||||||
|
pip install transformers torch
|
||||||
|
```
|
||||||
|
|
||||||
|
加载模型并生成文本:
|
||||||
|
```python
|
||||||
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
|
||||||
|
model = AutoModelForCausalLM.from_pretrained("XingChina/ChunMengDie-1.0-0.4B")
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained("XingChina/ChunMengDie-1.0-0.4B")
|
||||||
|
|
||||||
|
inputs = tokenizer("你好", return_tensors="pt")
|
||||||
|
outputs = model.generate(**inputs, max_new_tokens=128)
|
||||||
|
print(tokenizer.decode(outputs[0]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### 方式二:llama.cpp(GGUF 格式)
|
||||||
|
|
||||||
|
下载对应的 `.gguf` 文件后,使用 `llama.cpp` 进行推理:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# BF16 原始版本
|
||||||
|
./main -m chunmengdie_f16.gguf -p "你好" -n 128
|
||||||
|
|
||||||
|
# 8bit 量化版本
|
||||||
|
./main -m chunmengdie_q8_0.gguf -p "你好" -n 128
|
||||||
|
|
||||||
|
# 4bit 量化版本(低显存首选)
|
||||||
|
./main -m chunmengdie_q4_km.gguf -p "你好" -n 128
|
||||||
|
```
|
||||||
|
|
||||||
|
**使用 ollama 部署(可选)**:
|
||||||
|
```bash
|
||||||
|
ollama create chunmengdie -f ./Modelfile
|
||||||
|
ollama run chunmengdie
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 🛡️ 训练环境与权重状态声明
|
||||||
|
|
||||||
|
本模型基于 **NVIDIA RTX Pro 6000(96GB 显存)** 专业卡训练,采用针对大显存优化的高吞吐配置(大 Batch Size + 长序列)。
|
||||||
|
|
||||||
|
**发布权重说明**:
|
||||||
|
1. 本仓库仅提供 **纯推理权重**(HF 格式 + GGUF 格式),**不包含优化器(Optimizer)和调度器(Scheduler)状态**。
|
||||||
|
2. 由于训练环境依赖 **96GB 显存级的内存分配策略**,在消费级显卡(如 24GB 4090)或标准 A100(80GB)上**无法直接加载续训**。(由于我手滑删了最终调度器,你拿什么卡都不能续训,即使能继续训也很容易爆显存并且发挥不出最大实力,中间调度器不能和最终模型混用,没有上传)
|
||||||
|
3. 如需进行二次微调,建议使用本仓库的 BF16 权重作为基座,自行挂载新的优化器开始训练。
|
||||||
|
4. 项目根目录train[1].py已针对20~24G显存显卡进行微调,实测在4090上全程不会爆显存,请以最新版train[1].py为准
|
||||||
|
|
||||||
|
**我们鼓励二次创新,但尊重算力投入,请合理使用开源资源。**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 📚 训练数据声明
|
||||||
|
|
||||||
|
本模型在训练过程中参考/使用了以下公开数据集进行统计学习(点击可查看原始版权信息):
|
||||||
|
|
||||||
|
| 数据集 | 许可证 | 使用方式 | 版权归属 |
|
||||||
|
|--------|--------|----------|----------|
|
||||||
|
| [BelleGroup/train_0.5M_CN](https://huggingface.co/datasets/BelleGroup/train_0.5M_CN) | GPL-3.0 | 全量用于训练并通过统计学习从数据中提取对话模式 | 版权归 BelleGroup 所有 |
|
||||||
|
| [liumindmind/NekoQA-10K](https://huggingface.co/datasets/liumindmind/NekoQA-10K) | Apache-2.0 | 全量用于训练并通过统计学习从数据中提取对话模式 | 版权归 MindsRiverPonder 所有 |
|
||||||
|
| [cyberlangke/Nana-catgirl-dataset-110k](https://huggingface.co/datasets/cyberlangke/Nana-catgirl-dataset-110k) | MIT | 全量用于训练并通过统计学习从数据中提取对话模式 | 版权归 cyberlangke 所有 |
|
||||||
|
| [XingChina/ChunMengDie-1.0-User-Data](https://huggingface.co/datasets/XingChina/ChunMengDie-1.0-User-Data) | CC BY 4.0 | 本数据集为原创角色对话数据,全量用于训练并通过统计学习从数据中提取对话模式 | 版权归 XingChina 所有 |
|
||||||
|
|
||||||
|
**重要说明**:
|
||||||
|
1. 本模型权重**并非**上述任何数据集的“衍生代码副本”或“复制品”,模型仅通过梯度下降从数据中提取统计模式。
|
||||||
|
2. 本模型的输出内容由 AI 随机生成,**不代表对训练数据集的检索或分发**。
|
||||||
|
3. 在极低概率下(<0.1%),模型可能因统计噪声而产生与训练数据某一段落高度相似的文本,该现象属于概率性巧合,建议使用者对输出进行抽样审核。
|
||||||
|
4. 验证loss与训练loss均已降到1以下,但不能理解语义
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## ⚠️ 免责声明
|
||||||
|
|
||||||
|
本模型按“原样”(AS IS)提供,**不附带任何形式明示或默示的担保**。
|
||||||
|
|
||||||
|
使用者应自行承担使用本模型产生的一切后果,包括但不限于:
|
||||||
|
- 输出的准确性、安全性、合规性;
|
||||||
|
- 对第三方知识产权的潜在侵犯(若发生,属极小概率事件,我方不承担责任)。
|
||||||
|
|
||||||
|
建议在生产环境部署前配合敏感词过滤和输出重复检测模块。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 🚫 使用限制
|
||||||
|
|
||||||
|
- 本模型目前质量不佳,**不建议用于任何生产环境**。
|
||||||
|
- 商业使用需自行评估输出内容的合规性。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 📄 许可证
|
||||||
|
|
||||||
|
本模型权重采用 **BSD 3-Clause License** 发布。详见 [LICENSE](./LICENSE) 文件。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 📖 引用
|
||||||
|
|
||||||
|
如果你在研究中使用本模型,请引用:
|
||||||
|
|
||||||
|
```bibtex
|
||||||
|
@misc{ChunMengDie-1.0-0.4B,
|
||||||
|
author = {XingChina},
|
||||||
|
title = {ChunMengDie-1.0-0.4B: A Chinese Conversational AI Model},
|
||||||
|
year = {2026},
|
||||||
|
publisher = {Hugging Face},
|
||||||
|
url = {https://huggingface.co/XingChina/ChunMengDie-1.0-0.4B}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
**感谢你对 ChunMengDie 项目的关注!** 🎉
|
||||||
3
chunmengdie_f16[1].gguf
Normal file
3
chunmengdie_f16[1].gguf
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:2417f30ee0094b5d6002eaa0fba28e75f9cf679f9905b705b3106e72e22ee380
|
||||||
|
size 524994784
|
||||||
3
chunmengdie_q4_km[1].gguf
Normal file
3
chunmengdie_q4_km[1].gguf
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:b6250ec13225273bb9f0364f33a1304b34aa61e974ed095f3b7af8704161eba1
|
||||||
|
size 190016192
|
||||||
3
chunmengdie_q8_0[1].gguf
Normal file
3
chunmengdie_q8_0[1].gguf
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:a79faa1c2f10907d009aad5e540a0aff88f0e756fa2d665d493979d94e580ec8
|
||||||
|
size 288004384
|
||||||
35
config.json
Normal file
35
config.json
Normal file
@@ -0,0 +1,35 @@
|
|||||||
|
{
|
||||||
|
"activation_function": "gelu_new",
|
||||||
|
"add_cross_attention": false,
|
||||||
|
"architectures": [
|
||||||
|
"GPT2LMHeadModel"
|
||||||
|
],
|
||||||
|
"attn_pdrop": 0.3,
|
||||||
|
"bos_token_id": 50256,
|
||||||
|
"dtype": "float32",
|
||||||
|
"embd_pdrop": 0.3,
|
||||||
|
"eos_token_id": 50256,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"layer_norm_epsilon": 1e-05,
|
||||||
|
"model_type": "gpt2",
|
||||||
|
"n_ctx": 4096,
|
||||||
|
"n_embd": 1024,
|
||||||
|
"n_head": 16,
|
||||||
|
"n_inner": null,
|
||||||
|
"n_layer": 16,
|
||||||
|
"n_positions": 4096,
|
||||||
|
"pad_token_id": null,
|
||||||
|
"reorder_and_upcast_attn": false,
|
||||||
|
"resid_pdrop": 0.3,
|
||||||
|
"scale_attn_by_inverse_layer_idx": false,
|
||||||
|
"scale_attn_weights": true,
|
||||||
|
"summary_activation": null,
|
||||||
|
"summary_first_dropout": 0.1,
|
||||||
|
"summary_proj_to_labels": true,
|
||||||
|
"summary_type": "cls_index",
|
||||||
|
"summary_use_proj": true,
|
||||||
|
"tie_word_embeddings": true,
|
||||||
|
"transformers_version": "5.5.0",
|
||||||
|
"use_cache": false,
|
||||||
|
"vocab_size": 50257
|
||||||
|
}
|
||||||
9
generation_config.json
Normal file
9
generation_config.json
Normal file
@@ -0,0 +1,9 @@
|
|||||||
|
{
|
||||||
|
"_from_model_config": true,
|
||||||
|
"bos_token_id": 50256,
|
||||||
|
"eos_token_id": 50256,
|
||||||
|
"output_attentions": false,
|
||||||
|
"output_hidden_states": false,
|
||||||
|
"transformers_version": "5.5.0",
|
||||||
|
"use_cache": true
|
||||||
|
}
|
||||||
3
model.safetensors
Normal file
3
model.safetensors
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:5149d17d5ad92d4a8c16b34580d8e1c5ead36198e0ad2b810b725be690a83222
|
||||||
|
size 1028816568
|
||||||
250306
tokenizer.json
Normal file
250306
tokenizer.json
Normal file
File diff suppressed because it is too large
Load Diff
12
tokenizer_config.json
Normal file
12
tokenizer_config.json
Normal file
@@ -0,0 +1,12 @@
|
|||||||
|
{
|
||||||
|
"add_prefix_space": false,
|
||||||
|
"backend": "tokenizers",
|
||||||
|
"bos_token": "<|endoftext|>",
|
||||||
|
"eos_token": "<|endoftext|>",
|
||||||
|
"errors": "replace",
|
||||||
|
"is_local": false,
|
||||||
|
"model_max_length": 1024,
|
||||||
|
"pad_token": "<|endoftext|>",
|
||||||
|
"tokenizer_class": "GPT2Tokenizer",
|
||||||
|
"unk_token": "<|endoftext|>"
|
||||||
|
}
|
||||||
237
train[1].py
Normal file
237
train[1].py
Normal file
@@ -0,0 +1,237 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
# Copyright (c) 2026 XingChina
|
||||||
|
# SPDX-License-Identifier: BSD-3-Clause
|
||||||
|
# 本代码采用 BSD 3-Clause 许可证,详见项目根目录的 LICENSE 文件。
|
||||||
|
# 为4090显卡进行优化,保证普通显卡也能训练(而不是找大显存显卡)
|
||||||
|
# 已经拿4090测试过,全程没有崩溃
|
||||||
|
|
||||||
|
import os
|
||||||
|
import json
|
||||||
|
import signal
|
||||||
|
import sys
|
||||||
|
import torch
|
||||||
|
from transformers import (
|
||||||
|
GPT2Config,
|
||||||
|
GPT2LMHeadModel,
|
||||||
|
GPT2Tokenizer,
|
||||||
|
Trainer,
|
||||||
|
TrainingArguments,
|
||||||
|
DataCollatorForLanguageModeling,
|
||||||
|
)
|
||||||
|
from datasets import Dataset
|
||||||
|
from bitsandbytes.optim import Adam8bit
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logging.basicConfig(level=logging.INFO)
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ================== 模型配置 ==================
|
||||||
|
MODEL_CONFIG = {
|
||||||
|
"n_embd": 1024,
|
||||||
|
"n_layer": 16,
|
||||||
|
"n_head": 16,
|
||||||
|
"n_positions": 4096,
|
||||||
|
"vocab_size": 50257,
|
||||||
|
"n_ctx": 4096,
|
||||||
|
"resid_pdrop": 0.3,
|
||||||
|
"embd_pdrop": 0.3,
|
||||||
|
"attn_pdrop": 0.3,
|
||||||
|
}
|
||||||
|
|
||||||
|
# ================== 训练参数(为4090优化,保证普通显卡能复现) ==================
|
||||||
|
TRAIN_ARGS = {
|
||||||
|
"output_dir": "./checkpoints_chunmengdie_gpt2",
|
||||||
|
"per_device_train_batch_size": 2,
|
||||||
|
"gradient_accumulation_steps": 8,
|
||||||
|
"num_train_epochs": 1,
|
||||||
|
"learning_rate": 3e-4,
|
||||||
|
"weight_decay": 0.01,
|
||||||
|
"warmup_steps": 500,
|
||||||
|
"logging_steps": 10,
|
||||||
|
"save_steps": 500,
|
||||||
|
"save_total_limit": 3,
|
||||||
|
"bf16": True,
|
||||||
|
"report_to": "none",
|
||||||
|
"dataloader_num_workers": 0,
|
||||||
|
"eval_strategy": "steps",
|
||||||
|
"eval_steps": 500,
|
||||||
|
"load_best_model_at_end": True,
|
||||||
|
"metric_for_best_model": "eval_loss",
|
||||||
|
"greater_is_better": False,
|
||||||
|
"gradient_checkpointing": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
MAX_SEQ_LEN = 4096
|
||||||
|
CHECKPOINT_DIR = TRAIN_ARGS["output_dir"]
|
||||||
|
|
||||||
|
# ====== 数据文件列表(这里假设数据已经清理过,实际用的数据确实清理过) ======
|
||||||
|
# mengdie_train由于包含其他数据集,没有开源
|
||||||
|
DATA_FILES = [
|
||||||
|
"mengdie_train.json",
|
||||||
|
"Belle_open_source_0.5M.json",
|
||||||
|
]
|
||||||
|
|
||||||
|
# ================== 数据加载 ==================
|
||||||
|
def load_data(file_list):
|
||||||
|
texts = []
|
||||||
|
for fpath in file_list:
|
||||||
|
if not os.path.exists(fpath):
|
||||||
|
logger.warning(f"File not found: {fpath}, skipped")
|
||||||
|
continue
|
||||||
|
with open(fpath, 'r', encoding='utf-8') as f:
|
||||||
|
first_char = f.read(1)
|
||||||
|
f.seek(0)
|
||||||
|
if first_char == '[':
|
||||||
|
data = json.load(f)
|
||||||
|
logger.info(f"Loaded {len(data)} samples from {fpath} (JSON array)")
|
||||||
|
else:
|
||||||
|
data = []
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if line:
|
||||||
|
try:
|
||||||
|
data.append(json.loads(line))
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
pass
|
||||||
|
logger.info(f"Loaded {len(data)} samples from {fpath} (JSONL)")
|
||||||
|
|
||||||
|
for item in data:
|
||||||
|
try:
|
||||||
|
if "instruction" in item and "input" in item and "output" in item:
|
||||||
|
inst = item["instruction"].strip()
|
||||||
|
inp = item.get("input", "").strip()
|
||||||
|
out = item["output"].strip()
|
||||||
|
user_text = f"{inst}\n{inp}" if inp else inst
|
||||||
|
assistant_text = out
|
||||||
|
elif "input" in item and "output" in item:
|
||||||
|
user_text = item["input"].strip()
|
||||||
|
assistant_text = item["output"].strip()
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
texts.append(f"用户:{user_text}\n猫娘:{assistant_text}")
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
logger.info(f"Total texts: {len(texts)}")
|
||||||
|
return texts
|
||||||
|
|
||||||
|
# ================== Tokenization==================
|
||||||
|
def tokenize_function(examples):
|
||||||
|
tokenized = tokenizer(
|
||||||
|
examples["text"],
|
||||||
|
truncation=True,
|
||||||
|
max_length=MAX_SEQ_LEN - 1,
|
||||||
|
padding=False,
|
||||||
|
return_attention_mask=False,
|
||||||
|
)
|
||||||
|
tokenized["input_ids"] = [ids + [tokenizer.eos_token_id] for ids in tokenized["input_ids"]]
|
||||||
|
return tokenized
|
||||||
|
|
||||||
|
# ================== 紧急保存 ==================
|
||||||
|
def emergency_save(sig, frame):
|
||||||
|
logger.info("\n🛑 Saving emergency checkpoint...")
|
||||||
|
try:
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
model_to_save = model.cpu()
|
||||||
|
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
|
||||||
|
model_to_save.save_pretrained(os.path.join(CHECKPOINT_DIR, "emergency"))
|
||||||
|
tokenizer.save_pretrained(os.path.join(CHECKPOINT_DIR, "emergency"))
|
||||||
|
logger.info("✅ Emergency checkpoint saved.")
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Emergency save failed: {e}")
|
||||||
|
sys.exit(0)
|
||||||
|
|
||||||
|
# ================== 主程序 ==================
|
||||||
|
if __name__ == "__main__":
|
||||||
|
tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
|
||||||
|
tokenizer.pad_token = tokenizer.eos_token
|
||||||
|
|
||||||
|
raw_texts = load_data(DATA_FILES)
|
||||||
|
dataset = Dataset.from_dict({"text": raw_texts})
|
||||||
|
|
||||||
|
tokenized_dataset = dataset.map(
|
||||||
|
tokenize_function,
|
||||||
|
batched=True,
|
||||||
|
remove_columns=["text"],
|
||||||
|
num_proc=4,
|
||||||
|
load_from_cache_file=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 验证 EOS 是否添加成功
|
||||||
|
sample_ids = tokenized_dataset[0]["input_ids"]
|
||||||
|
logger.info(f"Sample last 5 tokens: {sample_ids[-5:]}")
|
||||||
|
logger.info(f"Last token is EOS: {sample_ids[-1] == tokenizer.eos_token_id}")
|
||||||
|
|
||||||
|
# 9:1 划分
|
||||||
|
total_len = len(tokenized_dataset)
|
||||||
|
train_size = int(0.9 * total_len)
|
||||||
|
eval_size = total_len - train_size
|
||||||
|
train_dataset, eval_dataset = torch.utils.data.random_split(
|
||||||
|
tokenized_dataset, [train_size, eval_size]
|
||||||
|
)
|
||||||
|
logger.info(f"Train: {train_size}, Eval: {eval_size}")
|
||||||
|
|
||||||
|
logger.info("Initializing GPT-2 model from scratch...")
|
||||||
|
config = GPT2Config(**MODEL_CONFIG)
|
||||||
|
model = GPT2LMHeadModel(config)
|
||||||
|
|
||||||
|
optimizer = Adam8bit(model.parameters(), lr=TRAIN_ARGS["learning_rate"])
|
||||||
|
|
||||||
|
training_args = TrainingArguments(
|
||||||
|
output_dir=CHECKPOINT_DIR,
|
||||||
|
per_device_train_batch_size=TRAIN_ARGS["per_device_train_batch_size"],
|
||||||
|
gradient_accumulation_steps=TRAIN_ARGS["gradient_accumulation_steps"],
|
||||||
|
num_train_epochs=TRAIN_ARGS["num_train_epochs"],
|
||||||
|
learning_rate=TRAIN_ARGS["learning_rate"],
|
||||||
|
weight_decay=TRAIN_ARGS["weight_decay"],
|
||||||
|
warmup_steps=TRAIN_ARGS["warmup_steps"],
|
||||||
|
logging_steps=TRAIN_ARGS["logging_steps"],
|
||||||
|
save_steps=TRAIN_ARGS["save_steps"],
|
||||||
|
save_total_limit=TRAIN_ARGS["save_total_limit"],
|
||||||
|
bf16=TRAIN_ARGS["bf16"],
|
||||||
|
report_to=TRAIN_ARGS["report_to"],
|
||||||
|
dataloader_num_workers=TRAIN_ARGS["dataloader_num_workers"],
|
||||||
|
optim="adamw_8bit",
|
||||||
|
eval_strategy=TRAIN_ARGS["eval_strategy"],
|
||||||
|
eval_steps=TRAIN_ARGS["eval_steps"],
|
||||||
|
load_best_model_at_end=TRAIN_ARGS["load_best_model_at_end"],
|
||||||
|
metric_for_best_model=TRAIN_ARGS["metric_for_best_model"],
|
||||||
|
greater_is_better=TRAIN_ARGS["greater_is_better"],
|
||||||
|
gradient_checkpointing=TRAIN_ARGS["gradient_checkpointing"],
|
||||||
|
)
|
||||||
|
|
||||||
|
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
|
||||||
|
|
||||||
|
trainer = Trainer(
|
||||||
|
model=model,
|
||||||
|
args=training_args,
|
||||||
|
train_dataset=train_dataset,
|
||||||
|
eval_dataset=eval_dataset,
|
||||||
|
data_collator=data_collator,
|
||||||
|
optimizers=(optimizer, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
signal.signal(signal.SIGINT, emergency_save)
|
||||||
|
|
||||||
|
# 自动恢复最新 checkpoint(中间checkpoint保留调度器,最终权重无调度器,这是手滑删了调度器的主要原因)
|
||||||
|
latest_checkpoint = None
|
||||||
|
if os.path.exists(CHECKPOINT_DIR):
|
||||||
|
checkpoints = [d for d in os.listdir(CHECKPOINT_DIR) if d.startswith("checkpoint-")]
|
||||||
|
if checkpoints:
|
||||||
|
latest_checkpoint = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))[-1]
|
||||||
|
latest_checkpoint = os.path.join(CHECKPOINT_DIR, latest_checkpoint)
|
||||||
|
logger.info(f"✅ Found checkpoint: {latest_checkpoint}, will resume from there.")
|
||||||
|
else:
|
||||||
|
if os.path.exists(os.path.join(CHECKPOINT_DIR, "emergency")):
|
||||||
|
latest_checkpoint = os.path.join(CHECKPOINT_DIR, "emergency")
|
||||||
|
logger.info(f"✅ Found emergency checkpoint, will resume from there.")
|
||||||
|
else:
|
||||||
|
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
|
||||||
|
|
||||||
|
logger.info("🚀 Starting training...")
|
||||||
|
trainer.train(resume_from_checkpoint=latest_checkpoint)
|
||||||
|
|
||||||
|
final_path = os.path.join(CHECKPOINT_DIR, "final_model")
|
||||||
|
model.save_pretrained(final_path)
|
||||||
|
tokenizer.save_pretrained(final_path)
|
||||||
|
logger.info(f"✅ Final model saved to {final_path}")
|
||||||
|
logger.info("📌 Convert to GGUF: python /path/to/llama.cpp/convert_hf_to_gguf.py ./checkpoints_chunmengdie_gpt2/final_model --outfile model.gguf --outtype q8_0")
|
||||||
Reference in New Issue
Block a user