import pandas as pd
import json

def convert_sst2_for_volcano():
    """将SST-2数据转换为火山方舟格式"""
    
    # 读取训练数据
    df = pd.read_csv("sst2_train.csv")
    
    volcano_data = []
    for _, row in df.iterrows():
        # 转换为对话格式
        volcano_data.append({
            "messages": [
                {
                    "role": "user",
                    "content": f"请判断以下电影评论的情感倾向，只回答'正面'或'负面'：\n\n评论：{row['text']}\n情感："
                },
                {
                    "role": "assistant", 
                    "content": "正面" if row['label'] == 1 else "负面"
                }
            ]
        })
    
    # 保存为JSONL格式
    with open("sst2_train_volcano.jsonl", "w", encoding="utf-8") as f:
        for item in volcano_data:
            f.write(json.dumps(item, ensure_ascii=False) + "\n")
    
    print(f"✅ 已转换 {len(volcano_data)} 条训练样本")
    
    # 转换验证数据用于评测
    df_val = pd.read_csv("sst2_validation.csv")
    eval_data = []
    
    for _, row in df_val.iterrows():
        eval_data.append({
            "query": f"请判断以下电影评论的情感倾向，只回答'正面'或'负面'：\n\n评论：{row['text']}\n情感：",
            "reference_response": "正面" if row['label'] == 1 else "负面"
        })
    
    # 保存评测数据
    eval_df = pd.DataFrame(eval_data)
    eval_df.to_excel("sst2_evaluation.xlsx", index=False)
    
    print(f"✅ 已转换 {len(eval_data)} 条评测样本")

if __name__ == "__main__":
    convert_sst2_for_volcano()
