import pandas as pd
import json, os

def convert_sst2_to_llamafactory():
    os.makedirs("data/sst2", exist_ok=True)

    def to_conv(df):
        records = []
        for _, row in df.iterrows():
            label = "正面" if int(row["label"]) == 1 else "负面"
            user = (
                "请判断以下电影评论的情感倾向，只回答'正面'或'负面'：\n\n"
                f"评论：{row['text']}\n情感："
            )
            records.append({
                "conversations": [
                    {"from": "user", "value": user},
                    {"from": "assistant", "value": label}
                ]
            })
        return records

    train = pd.read_csv("sst2_train.csv")
    dev = pd.read_csv("sst2_validation.csv")

    with open("data/sst2/train.json", "w", encoding="utf-8") as f:
        for r in to_conv(train):
            f.write(json.dumps(r, ensure_ascii=False) + "\n")

    with open("data/sst2/eval.json", "w", encoding="utf-8") as f:
        for r in to_conv(dev):
            f.write(json.dumps(r, ensure_ascii=False) + "\n")

    print("✅ 已生成 data/sst2/train.json 与 data/sst2/eval.json")

if __name__ == "__main__":
    convert_sst2_to_llamafactory()
