基于AST提取Python RL代码参数时数值异常的问题求助
解决RL参数提取中数值失真与Value列不符的问题
问题根源
- 原代码未处理
int(25e6)这类函数调用式赋值,直接输出AST结构字符串,导致Value列显示异常 - 未针对类型转换表达式(如科学计数法转int/float)做求值处理
- 依赖已废弃的
NameConstant节点(Python3.8+已合并到Constant) - 遗漏必要的
import ast导入语句
方案1:提取参数的实际运行值
适合需要获取参数实际数值的场景(如统计、计算),核心修改是增加对类型转换调用的求值逻辑:
import sys import ast from tabulate import tabulate def extract_parameters_with_values_from_file(file_path: str) -> dict: with open(file_path, 'r') as file: source_code = file.read() parameter_values = {} tree = ast.parse(source_code) for node in ast.walk(tree): if isinstance(node, ast.AnnAssign): parameter_name = node.target.id # 解析数据类型注解 if isinstance(node.annotation, ast.Name): data_type = node.annotation.id else: data_type = str(ast.dump(node.annotation)) # 解析参数值 parameter_value = None if isinstance(node.value, ast.Constant): parameter_value = node.value.value # 处理int/float类型转换(如int(25e6)) elif isinstance(node.value, ast.Call) and hasattr(node.value.func, 'id'): func_name = node.value.func.id if func_name in ('int', 'float') and len(node.value.args) == 1: arg_node = node.value.args[0] if isinstance(arg_node, ast.Constant): try: parameter_value = eval(f"{func_name}({arg_node.value})") except: parameter_value = ast.dump(node.value) else: parameter_value = ast.dump(node.value) # 尝试对其他简单表达式求值 else: try: parameter_value = eval(compile(ast.Expression(node.value), '', 'eval')) except: parameter_value = ast.dump(node.value) parameter_values[parameter_name] = (data_type, parameter_value) return parameter_values # 以下函数保持原逻辑不变 def read_parameters_from_txt(file_path: str) -> list: with open(file_path, 'r') as file: parameters = file.read().splitlines() return parameters def extract_parameters_with_values(parameter_names: list, python_file_path: str) -> list: parameter_values = extract_parameters_with_values_from_file(python_file_path) extracted_parameters = [] for parameter_name in parameter_names: if parameter_name in parameter_values: data_type, value = parameter_values[parameter_name] extracted_parameters.append([parameter_name, data_type, value]) return extracted_parameters if __name__ == "__main__": if len(sys.argv) != 3: sys.exit(1) parameter_txt_path = sys.argv[1] python_file_path = sys.argv[2] parameter_names = read_parameters_from_txt(parameter_txt_path) extracted_parameters = extract_parameters_with_values(parameter_names, python_file_path) if not extracted_parameters: print("No parameters found in the source code") else: print(tabulate(extracted_parameters, headers=['Parameter', 'Data Type', 'Value'], tablefmt="github"))
方案2:提取参数的原始代码表达式
适合需要保留代码中原始写法的场景(如文档生成),使用ast.get_source_segment直接获取代码片段:
import sys import ast from tabulate import tabulate def extract_parameters_with_values_from_file(file_path: str) -> dict: with open(file_path, 'r') as file: source_code = file.read() parameter_values = {} tree = ast.parse(source_code) for node in ast.walk(tree): if isinstance(node, ast.AnnAssign): parameter_name = node.target.id # 解析数据类型注解 if isinstance(node.annotation, ast.Name): data_type = node.annotation.id else: data_type = str(ast.dump(node.annotation)) # 获取原始表达式字符串 parameter_value = ast.get_source_segment(source_code, node.value) or ast.dump(node.value) parameter_values[parameter_name] = (data_type, parameter_value) return parameter_values # 其余函数与方案1一致,此处省略
注意事项
- 方案1中的
eval仅处理安全的字面量表达式,若代码存在复杂函数调用或变量引用,建议改用方案2 - 确保使用Python3.8+版本,避免
NameConstant的兼容性问题 - 原代码遗漏
import ast,必须添加才能正常运行
内容的提问来源于stack exchange,提问作者Brie MerryWeather
相关产品推荐
相关产品推荐

