# contrastive_text_classification **Repository Path**: letyygo_admin/contrastive_text_classification ## Basic Information - **Project Name**: contrastive_text_classification - **Description**: GLUT 大数据挖掘课程项目作业,30#706组内共同协作开发 - **Primary Language**: Unknown - **License**: Not specified - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 0 - **Created**: 2026-07-04 - **Last Updated**: 2026-07-08 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # 基于对比学习的文本分类 ## 一、环境准备 ```bash pip install -r requirements.txt ``` 把 `数据集说明.md` 里的 `dataset_pack/` 文件夹放到项目根目录下,结构应为: ``` dataset_pack/ ├── mr/{train,validation,test}.csv ├── ag_news/{train,validation,test}.csv └── trec/{train,validation,test}.csv ``` ## 二、快速开始(对应执行清单第 1-3 步) ```bash # 第1步:先在 MR 上跑通 baseline bash scripts/run_baseline.sh mr # 第2步:依次接入三种对比学习模块 bash scripts/run_simcse.sh mr bash scripts/run_consert.sh mr shuffle bash scripts/run_esimcse.sh mr # 第3步:确认都跑通后,搬到三个数据集上做主实验 bash scripts/run_all_main.sh ``` ## 三、消融实验(对应执行清单第 5 步) ```bash bash scripts/run_ablation.sh ``` 默认包含:对比损失权重 λ 扫描、温度系数 τ 扫描、低资源数据量扫描(10%/30%/100%)。 按需在 `run_ablation.sh` 里增删循环。 ## 四、结果汇总 ```bash python scripts/collect_results.py ``` 会在 `outputs/summary.csv` 生成一张包含 Accuracy / Macro-F1 / 每epoch耗时 / 最佳epoch 的汇总表, 可直接贴到报告里。 ## 五、核心设计说明 - **联合训练**:同一个 batch 里,分类损失 `cls_loss` 和对比损失 `cl_loss` 加权相加,`total = cls + λ * cl`。 - **正样本构造差异**(三种方法的核心区别): - SimCSE:同一输入两次过 encoder,靠 dropout 随机性产生差异。 - ConSERT:输入过 token shuffle 或 cutoff 增强后再过 encoder。 - ESimCSE:输入过 word repetition 增强,可选接动量编码器 + 负样本队列。 - **负样本**:默认都用 batch 内其它样本(in-batch negatives),ESimCSE 额外支持动量队列。 ## 六、项目目录结构说明 contrastive_text_classification/ ├── README.md # 项目说明(复制到你的仓库根目录即可) ├── requirements.txt # 依赖清单 ├── dataset_pack/ # 数据集放这里(mr/ag_news/trec 三个子目录) ├── src/ │ ├── utils.py # 随机种子、日志、计时工具 │ ├── augmentations.py # word repetition / token shuffle / cutoff 三种增强 │ ├── data_utils.py # 数据加载 + Collator(按 cl_type 自动生成增强样本) │ ├── model.py # ContrastiveClassifier(四种模式共用一套代码) │ ├── evaluate.py # Accuracy / Macro-F1 计算 │ └── train.py # 训练主脚本(记录耗时、最佳 epoch 等全部指标) ├── scripts/ │ ├── run_baseline.sh │ ├── run_simcse.sh │ ├── run_consert.sh │ ├── run_esimcse.sh │ ├── run_all_main.sh # 主实验:3数据集 × 4模型 一次跑完 │ ├── run_ablation.sh # 消融实验(λ、温度、数据量等) │ ├── collect_results.py # 汇总所有 json 结果成一张表 │ └── plot_results.py # 简单可视化 └── outputs/ # 训练日志、checkpoint、结果 json/csv 自动存这里