UCAS-EasyTranslate / scripts /visualize.py
lijn14
创建工程
c1a46f7
Raw
History Blame
2.45 kB
"""
可视化与分析脚本 — Person E 负责
功能:
1. 训练曲线绘制 (loss, BLEU, learning rate)
2. 实验结果对比图
3. 注意力权重可视化
4. 翻译样例展示
"""
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
def plot_training_curves(log_dir: str, output_path: str = "outputs/training_curves.png"):
"""
绘制训练曲线。
TODO [Person E]:
1. 从 TensorBoard 日志或 CSV 读取训练数据
2. 绘制子图:
- Train Loss vs Steps
- Val Loss vs Steps
- BLEU vs Epochs
- Learning Rate vs Steps
3. 保存图片
"""
raise NotImplementedError("TODO: Person E 实现 plot_training_curves")
def plot_experiment_comparison(results_dir: str, output_path: str = "outputs/experiment_comparison.png"):
"""
绘制实验对比图。
TODO [Person E]:
1. 读取所有实验的评估结果
2. 绘制柱状图: BLEU / COMET / chrF 对比
3. 绘制表格: 所有指标汇总
4. 保存图片
"""
raise NotImplementedError("TODO: Person E 实现 plot_experiment_comparison")
def visualize_attention(
model,
src_text: str,
tgt_text: str,
tokenizer,
output_path: str = "outputs/attention_map.png",
):
"""
注意力权重可视化。
TODO [Person E]:
1. 获取模型的 encoder self-attention 和 cross-attention 权重
2. 绘制热力图 (matplotlib / seaborn)
3. x 轴: 源语言 tokens, y 轴: 目标语言 tokens
4. 支持多头注意力的分别可视化和平均可视化
"""
raise NotImplementedError("TODO: Person E 实现 visualize_attention")
def generate_translation_examples(
evaluator,
test_pairs: list[tuple[str, str]],
output_path: str = "outputs/translation_examples.md",
):
"""
生成翻译样例展示。
TODO [Person E]:
1. 翻译测试样例
2. 生成 Markdown 格式的对比表格:
| 源文 (英) | 参考翻译 (中) | 模型翻译 (中) | BLEU |
3. 包含好的和差的翻译案例
4. 保存为 Markdown 文件
"""
raise NotImplementedError("TODO: Person E 实现 generate_translation_examples")
if __name__ == "__main__":
print("请指定要运行的可视化任务")
print(" python scripts/visualize.py --task training_curves --log_dir logs/")
print(" python scripts/visualize.py --task comparison --results_dir outputs/")