Text Generation
Chinese
English
roleplay
角色扮演
chinese
sft
conversational
creative-writing
digital-human
Qwen3-4B-Roleplay-Chinese / train_roleplay.py
Givenn's picture
Add roleplay SFT training script
573cc43 verified
Raw
History Blame
7.37 kB
#!/usr/bin/env python3
"""
Fine-tune Qwen3-4B for immersive Chinese roleplay (角色扮演).
Combines shibing624 roleplay-zh + ChatHaruhi-54K datasets.
Requirements: fast plot progression, strong immersion.
Usage:
pip install transformers trl torch datasets trackio accelerate peft
python train_roleplay.py
Hardware: a10g-largex2 (2x24GB GPU) recommended
Estimated time: ~4 hours for 3 epochs
"""
import os
import random
# Trackio monitoring setup
os.environ["TRACKIO_PROJECT"] = "qwen3-4b-roleplay"
from datasets import load_dataset, concatenate_datasets, Dataset
from trl import SFTTrainer, SFTConfig
# ============================================================
# Configuration
# ============================================================
MODEL_ID = "Qwen/Qwen3-4B"
OUTPUT_MODEL_ID = "Givenn/Qwen3-4B-Roleplay-Chinese"
MAX_SEQ_LENGTH = 4096
NUM_TRAIN_EPOCHS = 3
LEARNING_RATE = 2e-5
PER_DEVICE_BATCH_SIZE = 2
GRADIENT_ACCUMULATION_STEPS = 8 # effective batch = 16
# ============================================================
# 1. Load and prepare datasets
# ============================================================
print("=" * 60)
print("Loading datasets...")
print("=" * 60)
# --- Dataset 1: shibing624 roleplay-zh (ShareGPT format) ---
configs = [
"sharegpt_formatted_data-evol-gpt4",
"sharegpt_formatted_data-evol-gpt35",
"sharegpt_formatted_data-evol-male-gpt35",
"sharegpt_formatted_data-roleplay-chat-1k",
]
shibing_datasets = []
for cfg in configs:
ds = load_dataset(
"shibing624/roleplay-zh-sharegpt-gpt4-data",
name=cfg,
split="train",
)
shibing_datasets.append(ds)
print(f" Loaded shibing624/{cfg}: {len(ds)} samples")
# Convert shibing624 ShareGPT -> messages format
def convert_shibing_to_messages(example):
messages = []
if example.get("system_prompt") and example["system_prompt"].strip():
messages.append({
"role": "system",
"content": example["system_prompt"].strip()
})
for turn in example["conversations"]:
role = "user" if turn["from"] == "human" else "assistant"
messages.append({
"role": role,
"content": turn["value"]
})
return {"messages": messages}
converted_shibing = []
for ds in shibing_datasets:
converted = ds.map(
convert_shibing_to_messages,
remove_columns=ds.column_names,
num_proc=4,
)
converted_shibing.append(converted)
shibing_combined = concatenate_datasets(converted_shibing)
print(f"\nTotal shibing624 samples: {len(shibing_combined)}")
# --- Dataset 2: ChatHaruhi-54K (Chinese novel characters) ---
print("\nLoading ChatHaruhi-54K...")
haruhi_ds = load_dataset(
"silk-road/ChatHaruhi-54K-Role-Playing-Dialogue",
split="train",
)
print(f" Loaded ChatHaruhi-54K: {len(haruhi_ds)} samples")
def convert_haruhi_to_messages(example):
messages = []
agent_role = example.get("agent_role", "")
system_content = f'你现在扮演"{agent_role}"。请完全沉浸在角色中,用角色的语气、性格和说话方式来回应。保持角色一致性,推动剧情发展,营造沉浸感。'
messages.append({"role": "system", "content": system_content})
more_dialogues = example.get("more_dialogues", [])
if more_dialogues and len(more_dialogues) > 0:
for dialogue in more_dialogues:
if isinstance(dialogue, str) and ":" in dialogue:
parts = dialogue.split(":", 1)
if len(parts) == 2:
speaker = parts[0].strip()
content = parts[1].strip()
if speaker == agent_role:
messages.append({"role": "assistant", "content": content})
else:
messages.append({"role": "user", "content": f"({speaker}){content}"})
user_role = example.get("user_role", "")
user_question = example.get("user_question", "")
if user_role:
messages.append({"role": "user", "content": f"({user_role}){user_question}"})
else:
messages.append({"role": "user", "content": user_question})
agent_response = example.get("agent_response", "")
messages.append({"role": "assistant", "content": agent_response})
return {"messages": messages}
haruhi_converted = haruhi_ds.map(
convert_haruhi_to_messages,
remove_columns=haruhi_ds.column_names,
num_proc=4,
)
print(f" Converted ChatHaruhi: {len(haruhi_converted)} samples")
# ============================================================
# 2. Combine all datasets
# ============================================================
combined_dataset = concatenate_datasets([shibing_combined, haruhi_converted])
combined_dataset = combined_dataset.shuffle(seed=42)
print(f"\n{'=' * 60}")
print(f"Total combined dataset: {len(combined_dataset)} samples")
print(f"{'=' * 60}")
# Preview a sample
print("\n--- Sample data ---")
sample = combined_dataset[0]
for msg in sample["messages"][:3]:
print(f"[{msg['role']}]: {msg['content'][:100]}...")
print("---")
# ============================================================
# 3. Setup training with SFTTrainer
# ============================================================
print("\nInitializing SFTConfig...")
training_args = SFTConfig(
output_dir="./qwen3-4b-roleplay",
# Training hyperparameters
num_train_epochs=NUM_TRAIN_EPOCHS,
per_device_train_batch_size=PER_DEVICE_BATCH_SIZE,
gradient_accumulation_steps=GRADIENT_ACCUMULATION_STEPS,
learning_rate=LEARNING_RATE,
lr_scheduler_type="cosine",
warmup_steps=100,
weight_decay=0.01,
optim="adamw_torch_fused",
# Precision & memory
bf16=True,
gradient_checkpointing=True,
max_length=MAX_SEQ_LENGTH,
# Only train on assistant responses (loss masking)
completion_only_loss=True,
# Logging
logging_steps=5,
logging_first_step=True,
disable_tqdm=True,
report_to="trackio",
run_name="qwen3-4b-roleplay-zh",
# Saving & Hub
save_strategy="steps",
save_steps=500,
save_total_limit=3,
push_to_hub=True,
hub_model_id=OUTPUT_MODEL_ID,
hub_strategy="every_save",
# Dataset processing
dataset_num_proc=4,
# Seed
seed=42,
data_seed=42,
)
print("Initializing SFTTrainer...")
trainer = SFTTrainer(
model=MODEL_ID,
args=training_args,
train_dataset=combined_dataset,
)
# ============================================================
# 4. Train
# ============================================================
print(f"\n{'=' * 60}")
print("Starting training...")
print(f" Model: {MODEL_ID}")
print(f" Dataset size: {len(combined_dataset)}")
print(f" Epochs: {NUM_TRAIN_EPOCHS}")
print(f" Effective batch size: {PER_DEVICE_BATCH_SIZE * GRADIENT_ACCUMULATION_STEPS}")
print(f" Learning rate: {LEARNING_RATE}")
print(f" Max sequence length: {MAX_SEQ_LENGTH}")
print(f" Output: {OUTPUT_MODEL_ID}")
print(f"{'=' * 60}\n")
trainer.train()
# ============================================================
# 5. Save & push
# ============================================================
print("\nSaving final model...")
trainer.save_model()
trainer.push_to_hub()
print(f"\nModel pushed to: https://huggingface.co/{OUTPUT_MODEL_ID}")
print("Training complete!")