远程运行Captum Llama2归因代码时如何将生成图片保存至本地而非展示
远程运行Captum Llama2归因代码时如何将生成图片保存至本地而非展示
嗨,我之前在远程服务器上跑Captum的LLM归因代码时也碰到过一模一样的问题——SSH连过去没GUI,调用show()要么报错要么毫无反应,后来摸索出两个简单的解决办法,分享给你:
方法一:直接利用plot_token_attr的返回对象保存
Captum的plot_token_attr方法会返回Matplotlib的Figure对象,我们只需要把show=True改成show=False(避免尝试弹出图形窗口),然后用这个Figure对象的savefig方法直接保存到本地文件里就行。修改你的代码如下:
skip_tokens = [1] # skip the special token for the start of the text <s> inp = TextTokenInput( eval_prompt, tokenizer, skip_tokens=skip_tokens, ) target = "playing guitar, hiking, and spending time with his family." attr_res = llm_attr.attribute(inp, target=target, skip_tokens=skip_tokens) # 不显示图片,获取可视化的figure对象 fig = attr_res.plot_token_attr(show=False) # 保存到指定路径,这里以当前目录下的attribution_plot.png为例 fig.savefig("attribution_plot.png", bbox_inches='tight') # 可选:如果要避免内存占用,可以手动关闭figure import matplotlib.pyplot as plt plt.close(fig)
方法二:配置Matplotlib使用非交互式后端
如果你的代码里有多处可视化调用,不想逐个修改,也可以在代码最开头配置Matplotlib使用非交互式的后端(比如Agg),这样它就不会尝试启动图形界面,之后直接用Matplotlib的savefig方法保存即可:
# 在所有导入代码之前添加这两行,设置非交互式后端 import matplotlib matplotlib.use('Agg') import matplotlib.pyplot as plt # 下面是你原来的业务代码 skip_tokens = [1] # skip the special token for the start of the text <s> inp = TextTokenInput( eval_prompt, tokenizer, skip_tokens=skip_tokens, ) target = "playing guitar, hiking, and spending time with his family." attr_res = llm_attr.attribute(inp, target=target, skip_tokens=skip_tokens) # 关闭显示,然后保存图片 attr_res.plot_token_attr(show=False) plt.savefig("attribution_plot.png", bbox_inches='tight') plt.close()
小提示
- 保存路径可以用绝对路径,比如
/home/your_username/attribution_plots/plot1.png,记得提前创建好对应的目录,不然会抛出文件不存在的错误; bbox_inches='tight'这个参数很实用,能避免图片边缘的token文字被截断,建议加上;- 如果你之后需要在本地查看图片,可以用
scp命令把服务器上的图片下载到本地,比如scp your_username@server_ip:/path/to/attribution_plot.png ./local_path/
备注:内容来源于stack exchange,提问作者Jose Ramon
相关产品推荐
相关产品推荐

