trans_data_to / trans.py
cccxi's picture
Upload LoRA adapter folder
ed55dff verified
Raw
History Blame Contribute Delete
4.99 kB
import json
import os
from transformers import AutoTokenizer
# 1. 配置
# MODEL_PATH = "/home/at0842/ycl466704.ai13/.cache/huggingface/hub/models--openai--gpt-oss-20b/snapshots/6cee5e81ee83917806bbde320786a8fb61efebee"
# INPUT_JSONL = "multi_turn_zh_tw_function_mix500_turn_oss.jsonl"
# OUTPUT_JSONL = "multi_turn_zh_tw_function_mix500_turn_oss_gpt_oss_20b_pretokenized.jsonl"
# MAX_LENGTH = 14436
MODEL_PATH = "/home/at0842/ycl466704.ai13/.cache/huggingface/hub/models--openai--gpt-oss-20b/snapshots/6cee5e81ee83917806bbde320786a8fb61efebee"
INPUT_JSONL = "multi_turn_zh_tw_function_mix_v5_700_turn_oss.jsonl"
OUTPUT_JSONL = "multi_turn_zh_tw_function_mix_v5_700_turn_oss_gpt_oss_20b_pretokenized.jsonl"
MAX_LENGTH = 14436
if not os.path.exists(MODEL_PATH):
print(f"❌ 錯誤:找不到模型路徑 {MODEL_PATH}")
exit()
print(f"⏳ 正在初始化 GPT-OSS 20B Tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH, trust_remote_code=True)
# 取得結束 Token 的 ID (GPT-OSS 20B 專用)
RETURN_TOKEN_ID = tokenizer.convert_tokens_to_ids("<|return|>")
def process_entry(entry, index):
try:
messages = entry.get("messages", [])
raw_tools = entry.get("tools", None)
# --- 工具格式化 ---
formatted_tools = None
if raw_tools and isinstance(raw_tools, list):
formatted_tools = []
for t in raw_tools:
if isinstance(t, dict) and "function" in t:
formatted_tools.append(t)
else:
formatted_tools.append({
"type": "function",
"function": t
})
if not isinstance(messages, list) or len(messages) < 2:
return None, "對話輪次不足"
if messages[-1].get("role") != "assistant":
return None, f"最後一則不是 assistant (抓到的是: {messages[-1].get('role')})"
# A. 套用 Template (完整序列)
full_ids = tokenizer.apply_chat_template(
messages,
tools=formatted_tools,
tokenize=True,
add_generation_prompt=False,
truncation=True,
max_length=MAX_LENGTH
)
# --- 核心修正:手動補上遺失的 <|return|> ---
# 如果最後一個 ID 不是 <|return|>,我們手動補上
if full_ids[-1] != RETURN_TOKEN_ID:
full_ids.append(RETURN_TOKEN_ID)
# B. 計算 Context 長度 (提示部分)
context_ids = tokenizer.apply_chat_template(
messages[:-1],
tools=formatted_tools,
tokenize=True,
add_generation_prompt=True,
truncation=True,
max_length=MAX_LENGTH
)
start_idx = len(context_ids)
if start_idx >= len(full_ids):
return None, "內容超出 MAX_LENGTH 被截斷"
# C. 製作 Labels
# 遮蔽邏輯:非最後一輪 Assistant 全部設為 -100
labels = [-100] * len(full_ids)
for i in range(start_idx, len(full_ids)):
labels[i] = full_ids[i]
return {
"input_ids": full_ids,
"attention_mask": [1] * len(full_ids),
"labels": labels
}, None
except Exception as e:
return None, str(e)
# 2. 執行主程式
success_count = 0
drop_count = 0
print(f" 開始處理資料...")
with open(INPUT_JSONL, "r", encoding="utf-8") as f_in, \
open(OUTPUT_JSONL, "w", encoding="utf-8") as f_out:
for i, line in enumerate(f_in):
line = line.strip()
if not line: continue
try:
entry = json.loads(line)
result, error_msg = process_entry(entry, i)
if result:
f_out.write(json.dumps(result, ensure_ascii=False) + "\n")
success_count += 1
if success_count == 1:
print("\n" + "="*60)
print("【首筆資料切分預覽 - 修正後】")
loss_indices = [idx for idx, val in enumerate(result['labels']) if val != -100]
loss_tokens = [result['input_ids'][idx] for idx in loss_indices]
target_text = tokenizer.decode(loss_tokens)
print(f" Target (計算 Loss 部分): \n{target_text}")
last_tokens = tokenizer.convert_ids_to_tokens(result['input_ids'][-3:])
print(f" 整個序列末端三個 Token: {last_tokens}")
print("="*60 + "\n")
else:
drop_count += 1
except Exception as e:
drop_count += 1
print(f" 處理完成!")
print(f" 成功筆數: {success_count}")
print(f" 丟棄筆數: {drop_count}")
print(f" 結果檔案: {OUTPUT_JSONL}")