中文

用于PyTorch到JAX翻译的LLM学习错误情境

机器学习 2025-10-14 v1 人工智能

摘要

尽管近期大型语言模型(LLM)在主流语言的代码翻译方面取得了进展,但将PyTorch翻译到JAX仍颇具挑战。尽管这两个库都嵌入Python,但在核心设计、执行语义和生态系统成熟度方面存在差异;JAX较新且在公共代码中相对不足,PyTorch--JAX的平行语料库有限。现有评估方法的弱点进一步复杂化了跨框架基准测试。我们提出T2J,一个强化LLM基于PyTorch到JAX翻译的提示框架。我们的管道(i)组装两个PyTorch来源——来自TorchLeet的问题解决集(Aroori & Chien, 2025)和来自CodeParrot(Wolf et al., 2022)的GitHub派生集——并使用GPT-4o-mini生成初始JAX草稿;(ii)邀请两位专业开发人员迭代修复这些草稿直至功能等价,构建包含常见错误和补丁的精选固定-bug数据集;(iii)构建注入这些修复中结构化指导的增强提示,以引导轻量级LLM(如GPT-4o-mini)。我们还引入三个量身为PyTorch到JAX的指标:T2J CodeTrans Score、T2J FixCost Score(基于LLM的错误修复工作量估计)以及T2J Comparison Score(LLM作为评判员)。经验上,T2J将GPT-4o-mini性能提升最高可达10%(CodeBLEU),50%(T2J FixCost Score),1.33分(0-4比例的T2J CodeTrans Score),以及100%(T2J Comparison Score);此外,生成的代码比基线快至2.5倍。

关键词

引用

@article{arxiv.2510.09898,
  title  = {Learning Bug Context for PyTorch-to-JAX Translation with LLMs},
  author = {Hung Phan and Son Le Vu and Ali Jannesari},
  journal= {arXiv preprint arXiv:2510.09898},
  year   = {2025}
}