重构5,无io和无io依赖的模块已经全部重构完毕,接下来是重构剩余的模块,主要是io和依赖io的模块。
This commit is contained in:
@@ -0,0 +1,89 @@
|
||||
---
|
||||
name: fortran-extractor
|
||||
description: "[已完成] TLUSTY/SYNSPEC 拆分已完成。仅当用户明确请求'重新提取 Fortran'或'再次拆分 Fortran 文件'时触发。"
|
||||
---
|
||||
|
||||
# Fortran 代码提取器
|
||||
|
||||
从大型 Fortran 源文件中提取各个程序单元(SUBROUTINE、FUNCTION、PROGRAM、BLOCK DATA)到独立文件,并生成依赖分析报告。
|
||||
|
||||
## 快速参考
|
||||
|
||||
| 场景 | 命令 |
|
||||
|------|------|
|
||||
| 提取 Fortran 文件 | `python3 .claude/skills/fortran-extractor/scripts/extract_fortran.py <source.f> <output_dir>` |
|
||||
| 默认路径 | `tlusty/tlusty208.f` → `tlusty/extracted/` |
|
||||
|
||||
## 输出文件
|
||||
|
||||
提取完成后,输出目录包含:
|
||||
|
||||
| 文件 | 说明 |
|
||||
|------|------|
|
||||
| `*.f` | 各个提取的程序单元 |
|
||||
| `_SUMMARY.txt` | 提取摘要(单元数、类型统计) |
|
||||
| `_COMMON_ANALYSIS.txt` | COMMON 块依赖分析 |
|
||||
| `_PURE_UNITS.txt` | 无 COMMON 依赖的纯函数列表 |
|
||||
| `Makefile` | 编译配置(含正确标志) |
|
||||
|
||||
## 提取的单元类型
|
||||
|
||||
- `SUBROUTINE` - 子程序
|
||||
- `FUNCTION` - 函数
|
||||
- `PROGRAM` - 主程序(支持无名 PROGRAM)
|
||||
- `BLOCK DATA` - 数据块(支持无名 BLOCK DATA)
|
||||
|
||||
## Makefile 编译标志
|
||||
|
||||
生成的 Makefile 使用以下 gfortran 标志:
|
||||
|
||||
```makefile
|
||||
FFLAGS = -O3 -fno-automatic -mcmodel=large
|
||||
```
|
||||
|
||||
- `-mcmodel=large`: 支持大型 COMMON 数组(>2GB 地址空间)
|
||||
- `-fno-automatic`: 静态存储(旧 Fortran 兼容性)
|
||||
|
||||
**注意**: 不要使用 `-ffixed-line-length-none`,会破坏 73-80 列处理。
|
||||
|
||||
## COMMON 块分析
|
||||
|
||||
脚本自动分析每个单元的 COMMON 块依赖:
|
||||
|
||||
- **命名 COMMON**: `COMMON /NAME/ ...`
|
||||
- **空白 COMMON**: `COMMON varname`(不带斜杠)
|
||||
- **INCLUDE 依赖**: `INCLUDE 'XXX.FOR'`
|
||||
|
||||
### 纯函数识别
|
||||
|
||||
无 COMMON 依赖的单元被识别为"纯函数",可以独立测试和重构。
|
||||
|
||||
## 使用示例
|
||||
|
||||
```bash
|
||||
# 提取 TLUSTY
|
||||
python3 .claude/skills/fortran-extractor/scripts/extract_fortran.py tlusty/tlusty208.f tlusty/extracted/
|
||||
|
||||
# 提取 SYNSPEC
|
||||
python3 .claude/skills/fortran-extractor/scripts/extract_fortran.py synspec/synspec54.f synspec/extracted/
|
||||
|
||||
# 查看提取结果
|
||||
cat tlusty/extracted/_SUMMARY.txt
|
||||
cat tlusty/extracted/_PURE_UNITS.txt
|
||||
```
|
||||
|
||||
## 后续步骤
|
||||
|
||||
提取完成后,可以:
|
||||
|
||||
1. **编译验证**: `cd extracted && make`
|
||||
2. **依赖分析**: 使用 `fortran-analyzer` skill 分析函数调用依赖
|
||||
3. **重构**: 使用 `fortran-to-rust` skill 开始 Rust 重构
|
||||
|
||||
## 脚本位置
|
||||
|
||||
核心脚本位于 `.claude/skills/fortran-extractor/scripts/extract_fortran.py`,主要功能:
|
||||
|
||||
- `extract_units()`: 提取程序单元
|
||||
- `analyze_commons()`: 分析 COMMON 依赖
|
||||
- `generate_makefile()`: 生成编译配置
|
||||
@@ -0,0 +1,302 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
提取 synspec54.f 中的各个子程序/函数到独立文件
|
||||
"""
|
||||
import re
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
def extract_units(source_file, output_dir):
|
||||
"""提取 Fortran 程序单元到独立文件"""
|
||||
|
||||
with open(source_file, 'r') as f:
|
||||
content = f.read()
|
||||
lines = content.split('\n')
|
||||
|
||||
# 创建输出目录
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# 匹配程序单元开始的正则表达式
|
||||
# 注意: BLOCK DATA 和 PROGRAM 可以是无名的
|
||||
# 使用 \s* 允许名称前没有空格(无名情况)
|
||||
unit_pattern = re.compile(
|
||||
r'^\s*('
|
||||
r'SUBROUTINE\s+(\w+)|'
|
||||
r'FUNCTION\s+(\w+)|'
|
||||
r'PROGRAM\s*(\w*)|'
|
||||
r'BLOCK\s+DATA\s*(\w*)'
|
||||
r')',
|
||||
re.IGNORECASE
|
||||
)
|
||||
|
||||
# 找到所有单元的起始位置
|
||||
units = []
|
||||
for i, line in enumerate(lines):
|
||||
match = unit_pattern.match(line)
|
||||
if match:
|
||||
groups = match.groups()
|
||||
# groups: (整体匹配, SUBROUTINE名, FUNCTION名, PROGRAM名, BLOCK DATA名)
|
||||
|
||||
if groups[1]: # SUBROUTINE
|
||||
name, unit_type = groups[1], 'SUBROUTINE'
|
||||
elif groups[2]: # FUNCTION
|
||||
name, unit_type = groups[2], 'FUNCTION'
|
||||
elif groups[3]: # PROGRAM (非空)
|
||||
name, unit_type = groups[3], 'PROGRAM'
|
||||
elif groups[3] is not None: # PROGRAM (空字符串,无名)
|
||||
name, unit_type = None, 'PROGRAM'
|
||||
elif groups[4]: # BLOCK DATA (非空)
|
||||
name, unit_type = groups[4], 'BLOCK DATA'
|
||||
elif groups[4] is not None: # BLOCK DATA (空字符串,无名)
|
||||
name, unit_type = None, 'BLOCK DATA'
|
||||
else:
|
||||
name, unit_type = None, 'UNKNOWN'
|
||||
|
||||
# 处理无名单元
|
||||
if not name:
|
||||
name = f"_UNNAMED_{unit_type.replace(' ', '_')}_"
|
||||
|
||||
units.append((i, name.upper(), unit_type))
|
||||
|
||||
print(f"找到 {len(units)} 个程序单元")
|
||||
|
||||
# 提取每个单元
|
||||
extracted = []
|
||||
for idx, (start_line, name, unit_type) in enumerate(units):
|
||||
# 确定结束位置
|
||||
if idx + 1 < len(units):
|
||||
end_line = units[idx + 1][0]
|
||||
else:
|
||||
end_line = len(lines)
|
||||
|
||||
# 提取单元内容
|
||||
unit_lines = lines[start_line:end_line]
|
||||
|
||||
# 查找实际的 END 语句
|
||||
actual_end = end_line
|
||||
for i in range(len(unit_lines) - 1, -1, -1):
|
||||
if re.match(r'^\s*END\s*$', unit_lines[i], re.IGNORECASE):
|
||||
actual_end = start_line + i + 1
|
||||
break
|
||||
|
||||
unit_content = '\n'.join(lines[start_line:actual_end])
|
||||
|
||||
# 写入文件
|
||||
filename = f"{name.lower()}.f"
|
||||
filepath = os.path.join(output_dir, filename)
|
||||
|
||||
with open(filepath, 'w') as f:
|
||||
f.write(unit_content)
|
||||
if not unit_content.endswith('\n'):
|
||||
f.write('\n')
|
||||
|
||||
extracted.append({
|
||||
'name': name,
|
||||
'type': unit_type,
|
||||
'file': filename,
|
||||
'start': start_line + 1,
|
||||
'end': actual_end,
|
||||
'lines': actual_end - start_line
|
||||
})
|
||||
print(f" 提取: {name} ({unit_type}) -> {filename} ({actual_end - start_line} 行)")
|
||||
|
||||
# 生成摘要文件
|
||||
summary_path = os.path.join(output_dir, '_SUMMARY.txt')
|
||||
with open(summary_path, 'w') as f:
|
||||
f.write(f"SYNSPEC54.F 提取摘要\n")
|
||||
f.write(f"{'='*60}\n\n")
|
||||
f.write(f"源文件: {source_file}\n")
|
||||
f.write(f"总单元数: {len(extracted)}\n")
|
||||
f.write(f"总行数: {len(lines)}\n\n")
|
||||
|
||||
f.write(f"{'名称':<20} {'类型':<12} {'文件':<20} {'行数':>8}\n")
|
||||
f.write(f"{'-'*60}\n")
|
||||
for unit in extracted:
|
||||
f.write(f"{unit['name']:<20} {unit['type']:<12} {unit['file']:<20} {unit['lines']:>8}\n")
|
||||
|
||||
# 按类型统计
|
||||
types = {}
|
||||
for unit in extracted:
|
||||
types[unit['type']] = types.get(unit['type'], 0) + 1
|
||||
f.write(f"\n按类型统计:\n")
|
||||
for t, c in types.items():
|
||||
f.write(f" {t}: {c}\n")
|
||||
|
||||
print(f"\n摘要已保存到: {summary_path}")
|
||||
return extracted
|
||||
|
||||
def analyze_commons(output_dir):
|
||||
"""分析 COMMON 块依赖"""
|
||||
# 命名COMMON块: COMMON /NAME/ ...
|
||||
named_common_pattern = re.compile(r'COMMON\s*/\s*(\w+)\s*/', re.IGNORECASE)
|
||||
# 空白COMMON块: COMMON varname (不带斜杠)
|
||||
blank_common_pattern = re.compile(r'^\s*COMMON\s+[A-Z]', re.IGNORECASE | re.MULTILINE)
|
||||
include_pattern = re.compile(r'INCLUDE\s*[\'"]([^\'"]+)[\'"]', re.IGNORECASE)
|
||||
|
||||
commons = {}
|
||||
includes = {}
|
||||
|
||||
for filepath in Path(output_dir).glob('*.f'):
|
||||
if filepath.name.startswith('_'):
|
||||
continue
|
||||
|
||||
with open(filepath, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
unit_name = filepath.stem.upper()
|
||||
found_commons = named_common_pattern.findall(content)
|
||||
found_includes = include_pattern.findall(content)
|
||||
|
||||
# 检查空白COMMON块
|
||||
if blank_common_pattern.search(content):
|
||||
found_commons.append('BLANK') # 添加空白COMMON块标识
|
||||
|
||||
if found_commons:
|
||||
commons[unit_name] = list(set(found_commons))
|
||||
if found_includes:
|
||||
includes[unit_name] = list(set(found_includes))
|
||||
|
||||
# 写入 COMMON 分析
|
||||
common_path = os.path.join(output_dir, '_COMMON_ANALYSIS.txt')
|
||||
with open(common_path, 'w') as f:
|
||||
f.write("COMMON 块依赖分析\n")
|
||||
f.write(f"{'='*60}\n\n")
|
||||
|
||||
f.write("有 COMMON 依赖的单元:\n")
|
||||
f.write(f"{'-'*60}\n")
|
||||
for unit, common_list in sorted(commons.items()):
|
||||
f.write(f"{unit}: {', '.join(common_list)}\n")
|
||||
|
||||
f.write(f"\n共 {len(commons)} 个单元有 COMMON 依赖\n")
|
||||
f.write(f"共 {len([u for u in commons.values()])} 个 COMMON 块被引用\n")
|
||||
|
||||
# 找出所有唯一的 COMMON 块
|
||||
all_commons = set()
|
||||
for c in commons.values():
|
||||
all_commons.update(c)
|
||||
f.write(f"\n唯一的 COMMON 块: {sorted(all_commons)}\n")
|
||||
|
||||
f.write(f"\n\nINCLUDE 文件依赖:\n")
|
||||
f.write(f"{'-'*60}\n")
|
||||
for unit, inc_list in sorted(includes.items()):
|
||||
f.write(f"{unit}: {', '.join(inc_list)}\n")
|
||||
|
||||
print(f"COMMON 分析已保存到: {common_path}")
|
||||
|
||||
# 返回无 COMMON 依赖的纯函数
|
||||
pure_units = []
|
||||
for filepath in Path(output_dir).glob('*.f'):
|
||||
if filepath.name.startswith('_'):
|
||||
continue
|
||||
unit_name = filepath.stem.upper()
|
||||
if unit_name not in commons:
|
||||
pure_units.append(unit_name)
|
||||
|
||||
return pure_units, commons, includes
|
||||
|
||||
def generate_makefile(output_dir, extracted, source_file):
|
||||
"""生成 Makefile 用于编译所有提取的文件"""
|
||||
|
||||
# 根据源文件名确定程序名称
|
||||
source_name = os.path.basename(source_file).lower()
|
||||
if 'tlusty' in source_name:
|
||||
prog_name = 'tlusty'
|
||||
elif 'synspec' in source_name:
|
||||
prog_name = 'synspec'
|
||||
else:
|
||||
prog_name = os.path.splitext(os.path.basename(source_file))[0].lower()
|
||||
|
||||
makefile_path = os.path.join(output_dir, 'Makefile')
|
||||
with open(makefile_path, 'w') as f:
|
||||
f.write(f"# Makefile for {prog_name.upper()} extracted modules\n")
|
||||
f.write("# 使用大内存模型支持大型 COMMON 数组\n\n")
|
||||
|
||||
f.write("FC = gfortran\n")
|
||||
f.write("FFLAGS = -O3 -fno-automatic -mcmodel=large\n\n")
|
||||
|
||||
f.write("# 编译输出目录\n")
|
||||
f.write("BUILD_DIR = build\n\n")
|
||||
|
||||
f.write("# 目标可执行文件\n")
|
||||
f.write(f"MAIN = $(BUILD_DIR)/{prog_name}_extracted\n\n")
|
||||
|
||||
f.write("# 所有 .f 源文件\n")
|
||||
f.write("SRCS = $(wildcard *.f)\n\n")
|
||||
|
||||
f.write("# 目标文件(放在build目录)\n")
|
||||
f.write("OBJS = $(patsubst %.f,$(BUILD_DIR)/%.o,$(notdir $(SRCS)))\n\n")
|
||||
|
||||
f.write("# 默认目标\n")
|
||||
f.write("all: $(BUILD_DIR) $(MAIN)\n")
|
||||
f.write("\t@echo \"==========================================\"\n")
|
||||
f.write("\t@echo \"编译成功: $(MAIN)\"\n")
|
||||
f.write("\t@echo \"==========================================\"\n\n")
|
||||
|
||||
f.write("# 创建build目录\n")
|
||||
f.write("$(BUILD_DIR):\n")
|
||||
f.write("\tmkdir -p $(BUILD_DIR)\n\n")
|
||||
|
||||
f.write("# 链接所有目标文件\n")
|
||||
f.write("$(MAIN): $(OBJS)\n")
|
||||
f.write("\t$(FC) $(FFLAGS) -o $@ $(OBJS)\n\n")
|
||||
|
||||
f.write("# 编译规则\n")
|
||||
f.write("$(BUILD_DIR)/%.o: %.f | $(BUILD_DIR)\n")
|
||||
f.write("\t$(FC) $(FFLAGS) -c $< -o $@\n\n")
|
||||
|
||||
f.write("# 清理\n")
|
||||
f.write("clean:\n")
|
||||
f.write("\trm -rf $(BUILD_DIR)\n\n")
|
||||
|
||||
f.write("# 只编译不链接(检查语法)\n")
|
||||
f.write("compile-only: $(OBJS)\n")
|
||||
f.write("\t@echo \"所有文件编译完成(未链接)\"\n\n")
|
||||
|
||||
f.write("# 统计信息\n")
|
||||
f.write("stats:\n")
|
||||
f.write("\t@echo \"=== 编译统计 ===\"\n")
|
||||
f.write("\t@echo \"源文件数: $(words $(SRCS))\"\n")
|
||||
f.write("\t@echo \"目标文件数: $(words $(OBJS))\"\n")
|
||||
f.write("\t@wc -l *.f | tail -1\n\n")
|
||||
|
||||
f.write(".PHONY: all clean compile-only stats\n")
|
||||
|
||||
print(f"Makefile 已生成: {makefile_path}")
|
||||
|
||||
def main():
|
||||
if len(sys.argv) < 2:
|
||||
source_file = "/home/fmq/program/tlusty/tl208-s54/rust/synspec/synspec54.f"
|
||||
output_dir = "/home/fmq/program/tlusty/tl208-s54/rust/synspec/extracted"
|
||||
else:
|
||||
source_file = sys.argv[1]
|
||||
output_dir = sys.argv[2] if len(sys.argv) > 2 else "extracted"
|
||||
|
||||
print(f"源文件: {source_file}")
|
||||
print(f"输出目录: {output_dir}\n")
|
||||
|
||||
# 提取单元
|
||||
extracted = extract_units(source_file, output_dir)
|
||||
|
||||
# 分析 COMMON 依赖
|
||||
print("\n分析 COMMON 依赖...")
|
||||
pure_units, commons, includes = analyze_commons(output_dir)
|
||||
|
||||
print(f"\n无 COMMON 依赖的纯函数/子程序: {len(pure_units)} 个")
|
||||
for u in sorted(pure_units):
|
||||
print(f" {u}")
|
||||
|
||||
# 生成 Makefile
|
||||
generate_makefile(output_dir, extracted, source_file)
|
||||
|
||||
# 保存纯函数列表
|
||||
pure_path = os.path.join(output_dir, '_PURE_UNITS.txt')
|
||||
with open(pure_path, 'w') as f:
|
||||
f.write("无 COMMON 依赖的纯函数/子程序\n")
|
||||
f.write(f"{'='*40}\n\n")
|
||||
for u in sorted(pure_units):
|
||||
f.write(f"{u}\n")
|
||||
print(f"\n纯函数列表已保存到: {pure_path}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -0,0 +1,486 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
从 Fortran 源文件提取数组数据,生成 Rust data.rs
|
||||
|
||||
用法:
|
||||
python3 scripts/extract_fortran_data.py
|
||||
|
||||
输出: src/data.rs
|
||||
"""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def parse_fortran_file(filepath: Path, global_params: dict = None) -> list[dict]:
|
||||
"""解析单个 Fortran 文件中的数组
|
||||
|
||||
Args:
|
||||
filepath: Fortran 文件路径
|
||||
global_params: 从 include 文件中提取的全局参数表
|
||||
"""
|
||||
if global_params is None:
|
||||
global_params = {}
|
||||
|
||||
with open(filepath, 'r') as f:
|
||||
content = f.read()
|
||||
|
||||
arrays = {}
|
||||
|
||||
# 0. 预处理
|
||||
# 先清理每行的 Fortran 注释 (! 后面的内容)
|
||||
# 同时移除 Fortran 77 风格的注释行 (以 C 或 * 开头)
|
||||
lines = content.split('\n')
|
||||
cleaned_lines = []
|
||||
for line in lines:
|
||||
# Fortran 77 注释行 (第1列是 C, c, *, 或完全空行)
|
||||
if len(line) > 0 and line[0] in 'Cc*':
|
||||
continue # 跳过整行注释
|
||||
# Fortran 90 行内注释
|
||||
if '!' in line:
|
||||
line = line.split('!')[0]
|
||||
cleaned_lines.append(line)
|
||||
|
||||
# 合并 Fortran 续行
|
||||
# 方式1: 第6列是 *, +, &, 数字, 或字母 (固定格式)
|
||||
# 方式2: 行首 & (自由格式)
|
||||
merged_lines = []
|
||||
for line in cleaned_lines:
|
||||
# 自由格式续行: 行首 &
|
||||
if line.strip().startswith('&'):
|
||||
if merged_lines:
|
||||
merged_lines[-1] += ' ' + line.strip()[1:].strip()
|
||||
continue
|
||||
# 固定格式续行: 第6列非空
|
||||
if len(line) >= 6 and line[5] != ' ' and line[5] not in '\n\r\t':
|
||||
# 续行: 追加到上一行 (去掉前6列)
|
||||
if merged_lines:
|
||||
merged_lines[-1] += ' ' + line[6:].strip()
|
||||
else:
|
||||
merged_lines.append(line)
|
||||
content = '\n'.join(merged_lines)
|
||||
|
||||
# 1. 首先解析所有 parameter 语句,建立局部常量表
|
||||
# 合并全局参数和局部参数
|
||||
param_table = dict(global_params) # 复制全局参数
|
||||
param_pattern = r'parameter\s*\(([^)]+)\)'
|
||||
for match in re.finditer(param_pattern, content, re.IGNORECASE):
|
||||
params_str = match.group(1)
|
||||
for param in params_str.split(','):
|
||||
param = param.strip()
|
||||
if '=' in param:
|
||||
name, val = param.split('=', 1)
|
||||
name = name.strip().lower()
|
||||
val = val.strip().lower()
|
||||
# 尝试解析为整数或浮点数
|
||||
try:
|
||||
# 先尝试整数
|
||||
if '.' not in val and 'e' not in val and 'd' not in val:
|
||||
param_table[name] = int(val)
|
||||
else:
|
||||
# 浮点数,转换为整数(用于数组维度)
|
||||
val = val.replace('d', 'e')
|
||||
param_table[name] = int(float(val))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
# 2. 解析 dimension 语句
|
||||
# dimension p4a(22), p4b(10,28), adi(nni), ...
|
||||
# 注意: 使用 [ \t] 代替 \s 避免跨行匹配
|
||||
dim_pattern = r'dimension[ \t]+([a-z0-9_,() \t]+)'
|
||||
for match in re.finditer(dim_pattern, content, re.IGNORECASE):
|
||||
dim_str = match.group(1)
|
||||
# 清理 dim_str - 移除可能包含的下一个关键字
|
||||
for keyword in ['\nreal', '\ninteger', '\ncomplex', '\nlogical', '\ncharacter',
|
||||
'\ndimension', '\ndata', '\nparameter', '\nequivalence']:
|
||||
if keyword in dim_str.lower():
|
||||
dim_str = dim_str[:dim_str.lower().find(keyword)]
|
||||
break
|
||||
arr_pattern = r'(\w+)\s*\(([^)]+)\)'
|
||||
for arr_match in re.finditer(arr_pattern, dim_str):
|
||||
name = arr_match.group(1).lower()
|
||||
dims_str = arr_match.group(2)
|
||||
# 解析维度,支持常量和 parameter 变量
|
||||
dims = []
|
||||
valid = True
|
||||
for d in dims_str.split(','):
|
||||
d = d.strip().lower()
|
||||
if d in param_table:
|
||||
dims.append(param_table[d])
|
||||
else:
|
||||
try:
|
||||
dims.append(int(d))
|
||||
except ValueError:
|
||||
valid = False
|
||||
break
|
||||
if valid and dims:
|
||||
arrays[name] = {"name": name, "dims": dims, "data": None, "source": filepath.name}
|
||||
|
||||
# 2.5 解析类型声明中的数组
|
||||
# REAL frac(MR), INTEGER arr(10), REAL*4 arr(10), CHARACTER*10 str(5), etc.
|
||||
# 注意: 使用 [ \t] 代替 \s 避免跨行匹配
|
||||
type_decl_pattern = r'(real(?:\*[\d]+)?|integer(?:\*[\d]+)?|complex(?:\*[\d]+)?|logical(?:\*[\d]+)?|character(?:\*[\d]+)?)[ \t]+([a-z0-9_,() \t]+)'
|
||||
for match in re.finditer(type_decl_pattern, content, re.IGNORECASE):
|
||||
decl_str = match.group(2)
|
||||
# 清理 decl_str - 移除可能包含的下一个类型声明
|
||||
for keyword in ['\nreal', '\ninteger', '\ncomplex', '\nlogical', '\ncharacter',
|
||||
'\ndimension', '\ndata', '\nparameter', '\nequivalence']:
|
||||
if keyword in decl_str.lower():
|
||||
decl_str = decl_str[:decl_str.lower().find(keyword)]
|
||||
break
|
||||
|
||||
# 匹配变量名(维度)
|
||||
arr_pattern = r'(\w+)\s*\(([^)]+)\)'
|
||||
for arr_match in re.finditer(arr_pattern, decl_str):
|
||||
name = arr_match.group(1).lower()
|
||||
if name in arrays:
|
||||
continue # 已有定义
|
||||
dims_str = arr_match.group(2)
|
||||
# 解析维度
|
||||
dims = []
|
||||
valid = True
|
||||
for d in dims_str.split(','):
|
||||
d = d.strip().lower()
|
||||
if d in param_table:
|
||||
dims.append(param_table[d])
|
||||
else:
|
||||
try:
|
||||
dims.append(int(d))
|
||||
except ValueError:
|
||||
valid = False
|
||||
break
|
||||
if valid and dims:
|
||||
arrays[name] = {"name": name, "dims": dims, "data": None, "source": filepath.name}
|
||||
|
||||
# 3. 解析 data 语句 (支持多行和嵌套格式)
|
||||
# 格式1: data ((name(i,j),i=1,n),j=1,m)/values/
|
||||
# 格式2: data name /values/
|
||||
nested_data_pattern = r'data\s+\(\(\s*(\w+)\s*\([^)]+\)\s*,\s*\w+\s*=\s*\d+\s*,\s*\d+\s*\)\s*,\s*\w+\s*=\s*\d+\s*,\s*\d+\s*\)\s*/\s*([^/]+)\s*/'
|
||||
for match in re.finditer(nested_data_pattern, content, re.IGNORECASE | re.DOTALL):
|
||||
name = match.group(1).lower()
|
||||
data_str = match.group(2)
|
||||
|
||||
if name not in arrays:
|
||||
arrays[name] = {"name": name, "dims": [], "data": [], "source": filepath.name}
|
||||
|
||||
values = parse_data_values(data_str)
|
||||
# 合并多个 DATA 语句的值
|
||||
if arrays[name]["data"] is None:
|
||||
arrays[name]["data"] = values
|
||||
else:
|
||||
arrays[name]["data"].extend(values)
|
||||
|
||||
# 简单格式: data name /values/
|
||||
simple_data_pattern = r'data\s+(\w+)\s*/\s*([^/]+)\s*/'
|
||||
for match in re.finditer(simple_data_pattern, content, re.IGNORECASE | re.DOTALL):
|
||||
name = match.group(1).lower()
|
||||
data_str = match.group(2)
|
||||
|
||||
if name not in arrays:
|
||||
arrays[name] = {"name": name, "dims": [], "data": [], "source": filepath.name}
|
||||
|
||||
values = parse_data_values(data_str)
|
||||
# 合并多个 DATA 语句的值
|
||||
if arrays[name]["data"] is None:
|
||||
arrays[name]["data"] = values
|
||||
else:
|
||||
arrays[name]["data"].extend(values)
|
||||
|
||||
# 4. 处理 parameter 语句中的标量常量(用于导出)
|
||||
for match in re.finditer(param_pattern, content, re.IGNORECASE):
|
||||
params_str = match.group(1)
|
||||
for param in params_str.split(','):
|
||||
param = param.strip()
|
||||
if '=' in param:
|
||||
name, val = param.split('=', 1)
|
||||
name = name.strip().lower()
|
||||
val = val.strip().lower()
|
||||
if name not in arrays:
|
||||
try:
|
||||
val = val.replace('d', 'e')
|
||||
arrays[name] = {"name": name, "dims": [], "data": [float(val)], "source": filepath.name, "is_param": True}
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return list(arrays.values())
|
||||
|
||||
|
||||
def parse_data_values(data_str: str) -> list[float]:
|
||||
"""解析 DATA 语句中的数值,处理重复语法如 7*1.387"""
|
||||
values = []
|
||||
|
||||
# 清理
|
||||
lines = data_str.split('\n')
|
||||
cleaned_lines = []
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
if line.startswith('*'):
|
||||
line = line[1:].strip()
|
||||
cleaned_lines.append(line)
|
||||
data_str = ' '.join(cleaned_lines)
|
||||
|
||||
for part in data_str.split(','):
|
||||
part = part.strip()
|
||||
if not part:
|
||||
continue
|
||||
|
||||
# 移除末尾的 / (DATA 语句结束符)
|
||||
part = part.rstrip('/')
|
||||
|
||||
# 处理重复语法: "7*1.387"
|
||||
if '*' in part and not part.startswith('-'):
|
||||
match = re.match(r'(\d+)\s*\*\s*(-?[\d.]+)', part)
|
||||
if match:
|
||||
count = int(match.group(1))
|
||||
val = float(match.group(2))
|
||||
values.extend([val] * count)
|
||||
continue
|
||||
|
||||
try:
|
||||
# 处理 Fortran 科学计数法
|
||||
# 处理 "- 14.2" 这种中间有空格的负数
|
||||
val = part.replace('d', 'e').replace('D', 'e')
|
||||
# 移除负号和数字之间的空格
|
||||
val = re.sub(r'-\s+(\d)', r'-\1', val)
|
||||
# 移除科学计数法中的多余空格 (如 "1.48 e-2" -> "1.48e-2")
|
||||
val = re.sub(r'(\d)\s+([eEdD])', r'\1\2', val)
|
||||
values.append(float(val))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return values
|
||||
|
||||
|
||||
def generate_data_rs(all_arrays: dict[str, list[dict]]) -> str:
|
||||
"""生成 src/data.rs 内容"""
|
||||
lines = []
|
||||
lines.append("//! Fortran 数据数组自动导出")
|
||||
lines.append("//!")
|
||||
lines.append("//! 由 extract_fortran_data.py 自动生成,请勿手动修改")
|
||||
lines.append("")
|
||||
|
||||
# 收集已使用的名称,避免重复
|
||||
used_names = set()
|
||||
|
||||
# 按源文件分组
|
||||
for source, arrays in sorted(all_arrays.items()):
|
||||
# 过滤有数据的数组
|
||||
valid_arrays = [a for a in arrays if a.get("data") and len(a["data"]) > 0]
|
||||
if not valid_arrays:
|
||||
continue
|
||||
|
||||
lines.append(f"// ========== {source} ==========")
|
||||
lines.append("")
|
||||
|
||||
for arr in valid_arrays:
|
||||
base_name = arr["name"].upper()
|
||||
# 清理名称中的特殊字符
|
||||
base_name = re.sub(r'[^A-Z0-9_]', '', base_name)
|
||||
|
||||
# 所有变量都添加文件名前缀,避免命名冲突
|
||||
prefix = Path(source).stem.upper()[:8] # 取文件名前8个字符
|
||||
name = f"{prefix}_{base_name}"
|
||||
|
||||
# 如果加上前缀后仍有冲突,添加序号
|
||||
if name in used_names:
|
||||
counter = 1
|
||||
while f"{name}_{counter}" in used_names:
|
||||
counter += 1
|
||||
name = f"{name}_{counter}"
|
||||
|
||||
used_names.add(name)
|
||||
|
||||
dims = arr["dims"]
|
||||
data = arr["data"]
|
||||
|
||||
if not data:
|
||||
continue
|
||||
|
||||
total_size = len(data)
|
||||
|
||||
# 跳过单值 parameter
|
||||
if arr.get("is_param") and total_size == 1:
|
||||
lines.append(f"/// {arr['name']} (from {source})")
|
||||
lines.append(f"pub const {name}: f64 = {data[0]};")
|
||||
lines.append("")
|
||||
continue
|
||||
|
||||
if len(dims) == 0:
|
||||
# 未知维度,用 Vec 格式输出以便检查
|
||||
lines.append(f"/// {arr['name']} (from {source}, 未知维度,共 {len(data)} 个值)")
|
||||
lines.append(f"pub const {name}: [f64; {len(data)}] = [")
|
||||
for i, val in enumerate(data):
|
||||
if i % 10 == 0:
|
||||
lines.append(" ")
|
||||
lines[-1] += f"{val},"
|
||||
lines.append("];")
|
||||
lines.append("")
|
||||
|
||||
elif len(dims) == 1:
|
||||
# 1D 数组 - 检查数据量是否匹配
|
||||
expected_size = dims[0]
|
||||
if len(data) != expected_size:
|
||||
print(f"警告: {name} 期望 {expected_size} 个值,实际 {len(data)} 个,跳过")
|
||||
continue
|
||||
|
||||
lines.append(f"/// {arr['name']}({dims[0]}) from {source}")
|
||||
lines.append(f"pub const {name}: [f64; {dims[0]}] = [")
|
||||
for i, val in enumerate(data):
|
||||
if i % 10 == 0:
|
||||
lines.append(" ")
|
||||
lines[-1] += f"{val},"
|
||||
lines.append("];")
|
||||
lines.append("")
|
||||
|
||||
elif len(dims) == 2:
|
||||
# 2D 数组 - 直接转换为 Rust 行优先格式
|
||||
nj, ni = dims[0], dims[1]
|
||||
expected_size = nj * ni
|
||||
|
||||
if len(data) < expected_size:
|
||||
print(f"警告: {name} 期望 {expected_size} 个值,实际 {len(data)} 个,跳过")
|
||||
continue
|
||||
|
||||
lines.append(f"/// {arr['name']}({nj}, {ni}) from {source}")
|
||||
lines.append(f"/// 已转换为 Rust 行优先格式")
|
||||
lines.append(f"pub const {name}: [[f64; {ni}]; {nj}] = [")
|
||||
|
||||
# 列优先 → 行优先 转换
|
||||
for j in range(nj):
|
||||
row = []
|
||||
for i in range(ni):
|
||||
idx = j + i * nj # Fortran 列优先索引
|
||||
row.append(str(data[idx]))
|
||||
lines.append(f" [{','.join(row)}],")
|
||||
|
||||
lines.append("];")
|
||||
lines.append("")
|
||||
|
||||
elif len(dims) == 3:
|
||||
# 3D 数组 - 转换为 Rust 格式 (使用扁平数组 + 索引计算)
|
||||
nk, nj, ni = dims[0], dims[1], dims[2]
|
||||
expected_size = nk * nj * ni
|
||||
|
||||
if len(data) < expected_size:
|
||||
print(f"警告: {name} 期望 {expected_size} 个值,实际 {len(data)} 个,跳过")
|
||||
continue
|
||||
|
||||
lines.append(f"/// {arr['name']}({nk}, {nj}, {ni}) from {source}")
|
||||
lines.append(f"/// Fortran 列优先存储,访问方式: data[k + nk*(j + nj*i)]")
|
||||
lines.append(f"pub const {name}: [f64; {expected_size}] = [")
|
||||
|
||||
for i, val in enumerate(data[:expected_size]):
|
||||
if i % 5 == 0:
|
||||
lines.append(" ")
|
||||
lines[-1] += f"{val},"
|
||||
|
||||
lines.append("];")
|
||||
lines.append("")
|
||||
|
||||
# 不再需要转换函数和 getter,2D 数组直接生成为 const
|
||||
|
||||
return '\n'.join(lines)
|
||||
|
||||
|
||||
def parse_include_files(extracted_dir: Path) -> dict:
|
||||
"""解析 .FOR include 文件中的全局参数"""
|
||||
global_params = {}
|
||||
|
||||
# 扫描 .FOR 文件
|
||||
for for_file in extracted_dir.glob("*.FOR"):
|
||||
try:
|
||||
content = for_file.read_text()
|
||||
except:
|
||||
# 也尝试 tlusty/ 根目录
|
||||
for_file = Path("tlusty") / for_file.name
|
||||
if for_file.exists():
|
||||
content = for_file.read_text()
|
||||
else:
|
||||
continue
|
||||
|
||||
# 清理注释
|
||||
lines = []
|
||||
for line in content.split('\n'):
|
||||
if '!' in line:
|
||||
line = line.split('!')[0]
|
||||
lines.append(line)
|
||||
content = '\n'.join(lines)
|
||||
|
||||
# 解析 parameter 语句
|
||||
param_pattern = r'parameter\s*\(([^)]+)\)'
|
||||
for match in re.finditer(param_pattern, content, re.IGNORECASE):
|
||||
params_str = match.group(1)
|
||||
for param in params_str.split(','):
|
||||
param = param.strip()
|
||||
if '=' in param:
|
||||
name, val = param.split('=', 1)
|
||||
name = name.strip().lower()
|
||||
val = val.strip().lower()
|
||||
try:
|
||||
if '.' not in val and 'e' not in val and 'd' not in val:
|
||||
global_params[name] = int(val)
|
||||
else:
|
||||
val = val.replace('d', 'e')
|
||||
global_params[name] = int(float(val))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return global_params
|
||||
|
||||
|
||||
def main():
|
||||
# 扫描 tlusty/extracted 目录
|
||||
extracted_dir = Path("tlusty/extracted")
|
||||
|
||||
if not extracted_dir.exists():
|
||||
print(f"错误: 目录不存在: {extracted_dir}")
|
||||
return
|
||||
|
||||
# 首先解析 include 文件中的全局参数
|
||||
global_params = parse_include_files(extracted_dir)
|
||||
print(f"从 .FOR include 文件中提取了 {len(global_params)} 个全局参数")
|
||||
|
||||
all_arrays = {}
|
||||
|
||||
# 扫描所有 .f 文件
|
||||
for fortran_file in sorted(extracted_dir.glob("*.f")):
|
||||
arrays = parse_fortran_file(fortran_file, global_params)
|
||||
if arrays:
|
||||
all_arrays[fortran_file.name] = arrays
|
||||
print(f"解析: {fortran_file.name} -> {len(arrays)} 个数组")
|
||||
|
||||
# 统计
|
||||
total_arrays = sum(len(arrs) for arrs in all_arrays.values())
|
||||
arrays_with_data = sum(
|
||||
1 for arrs in all_arrays.values()
|
||||
for a in arrs if a.get("data") and len(a["data"]) > 0
|
||||
)
|
||||
arrays_2d = sum(
|
||||
1 for arrs in all_arrays.values()
|
||||
for a in arrs if len(a.get("dims", [])) == 2 and a.get("data")
|
||||
)
|
||||
|
||||
print()
|
||||
print("=" * 60)
|
||||
print(f"总计: {total_arrays} 个数组, {arrays_with_data} 个有数据, {arrays_2d} 个 2D 数组")
|
||||
print("=" * 60)
|
||||
|
||||
# 生成 data.rs
|
||||
output_path = Path("src/data.rs")
|
||||
rust_code = generate_data_rs(all_arrays)
|
||||
|
||||
with open(output_path, 'w') as f:
|
||||
f.write(rust_code)
|
||||
|
||||
print(f"已生成: {output_path}")
|
||||
print()
|
||||
print("在 lib.rs 或 main.rs 中添加:")
|
||||
print(" pub mod data;")
|
||||
print()
|
||||
print("使用方法:")
|
||||
print(" use crate::data::{TT, PN, get_p4b};")
|
||||
print(" let p4b = get_p4b(); // 自动初始化并返回 2D 数组")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user