You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何从字符串中提取Tensor?求TensorFlow/PyTorch方法或Python实现

从字符串提取PyTorch Tensor的方法

首先明确:PyTorch或者TensorFlow目前都没有提供类似ast.literal_eval专门用于解析tensor(...)格式字符串的官方函数,得自己实现处理逻辑。

你初步想到的方法其实已经很靠谱了——先剥离外层的tensor(和),再用ast.literal_eval解析成列表,最后转成Tensor。这个思路安全且高效,毕竟ast.literal_eval只会解析合法的Python字面量,不会执行恶意代码,比直接用eval()稳妥得多。

不过可以稍微优化一下实现,让它更健壮:

  • 用正则表达式匹配括号内的数组部分,避免多次replace可能带来的意外(比如字符串里存在多余(或)的情况)
  • 直接用torch.tensor()替代torch.from_numpy(np.array(l)),写法更简洁

优化后的代码示例:

import ast
import re
import torch

tensor_str = "tensor([-1.6975e+00, 1.7556e-02, -2.4441e+00, -2.3994e+00, -6.2069e-01])"

# 用正则提取括号内的数组字符串
match = re.search(r'tensor\((.*)\)', tensor_str)
if match:
    list_str = match.group(1)
    num_list = ast.literal_eval(list_str)
    tensor = torch.tensor(num_list)
    print(tensor)
else:
    raise ValueError("Invalid tensor string format")

如果你的字符串还包含设备信息(比如tensor([1,2,3], device='cuda:0')),可以再扩展正则逻辑,把设备参数也提取出来,创建Tensor时直接指定设备:

import ast
import re
import torch

tensor_str = "tensor([-1.6975e+00, 1.7556e-02], device='cuda:0')"

match = re.search(r'tensor\((.*?)(?:, device=(.*?))?\)', tensor_str)
if match:
    list_str = match.group(1)
    device_str = match.group(2)
    num_list = ast.literal_eval(list_str)
    device = torch.device(device_str.strip("'\"")) if device_str else None
    tensor = torch.tensor(num_list, device=device)
    print(tensor)
else:
    raise ValueError("Invalid tensor string format")

总结来说,你的初始方案已经是Python风格的安全实现,优化版用正则让它的适应性更强,不管是基础格式还是带设备参数的字符串都能处理。

内容的提问来源于stack exchange,提问作者Belkacem Thiziri

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.08 12:58:15