初始化项目,由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