如何从字符串中提取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
相关产品推荐
相关产品推荐

