cwenzi/neuroflow-cpp
1
1import re
2import json
3import argparse
4import logging
5import os
6
7logger = logging.getLogger(__name__)
8
9MODEL_CONFIG_FIELDS = [
10 ("input_dim", 512, "int", "输入维度(d_model)", "hidden_dim"),
11 ("hidden_dim", 256, "int", "隐藏层维度", "hidden_dim"),
12 ("output_dim", 10, "int", "输出维度(分类数)", "num_labels"),
13 ("memory_dim", 128, "int", "记忆维度(d_mem)", None),
14 ("memory_slots", 64, "int", "记忆槽数(MEM_SLOTS)", None),
15 ("num_layers", 2, "int", "网络层数", "num_layers"),
16 ("num_associations", 8, "int", "DMN关联头数", None),
17 ("use_quantization", False, "bool", "是否启用量化", None),
18 ("use_mla", False, "bool", "是否启用MLA注意力", "use_mla"),
19 ("mla_latent_dim", 32, "int", "MLA潜在维度", None),
20 ("use_causal_lm", False, "bool", "是否启用因果LM", None),
21 ("vocab_size", 5000, "int", "词表大小(VOCAB_SIZE)", "vocab_size"),
22 ("max_seq_len", 128, "int", "最大序列长度", "max_seq_len"),
23 ("causal_window_size", 32, "int", "因果窗口大小", None),
24 ("sae_k", 64, "int", "SAE稀疏度", None),
25 ("ntm_memory_slots", 16, "int", "NTM记忆槽数", None),
26]
27
28CAUSAL_LM_FIELDS = [
29 ("vocab_size", 5000, "int", "词表大小", "vocab_size"),
30 ("d_model", 256, "int", "模型维度", "hidden_dim"),
31 ("max_seq_len", 128, "int", "最大序列长度", "max_seq_len"),
32 ("causal_window_size", 32, "int", "因果窗口大小", None),
33 ("sae_k", 64, "int", "SAE稀疏度", None),
34 ("ntm_memory_slots", 16, "int", "NTM记忆槽数", None),
35 ("use_mla", True, "bool", "启用MLA", None),
36 ("mla_latent_dim", 32, "int", "MLA潜在维度", None),
37 ("mla_n_heads", 8, "int", "MLA头数", None),
38 ("mla_max_cache_len", 4096, "int", "MLA最大缓存长度", None),
39 ("use_quantization", False, "bool", "启用量化", None),
40]
41
42GENERATE_FIELDS = [
43 ("max_new_tokens", 50, "int", "最大生成token数", None),
44 ("temperature", 1.0, "float", "采样温度(0.0,2.0]", None),
45 ("top_k", 40, "int", "Top-K采样K值", None),
46 ("top_p", 0.9, "float", "Top-P采样P值[0.0,1.0]", None),
47 ("repetition_penalty", 1.0, "float", "重复惩罚[1.0,2.0]", None),
48 ("punct_penalty", 0.0, "float", "标点惩罚", None),
49 ("random_seed", 0, "int", "随机种子(0=随机)", None),
50 ("strategy", "top_k", "str", "采样策略(greedy/top_k/top_p/top_k_top_p)", None),
51 ("eos_id", 3, "int", "EOS token ID", None),
52]
53
54
55def parse_config_from_hpp(hpp_path, struct_name):
56 fields = []
57 try:
58 with open(hpp_path, "r", encoding="utf-8", errors="ignore") as f:
59 content = f.read()
60 pattern = rf'struct\s+{struct_name}\s*\{{([^}}]+)\}}'
61 match = re.search(pattern, content, re.DOTALL)
62 if not match:
63 logger.warning(f"未找到结构体 {struct_name},使用默认字段")
64 return None
65 body = match.group(1)
66 for line in body.strip().split("\n"):
67 line = line.strip()
68 if not line or line.startswith("//"):
69 continue
70 m = re.match(r'(size_t|bool|float|std::string|int)\s+(\w+)\s*=\s*([^;]+);', line)
71 if m:
72 ctype, name, default = m.groups()
73 default = default.strip().rstrip('f')
74 if ctype == "size_t" or ctype == "int":
75 default = int(default)
76 elif ctype == "float":
77 default = float(default)
78 elif ctype == "bool":
79 default = default == "true"
80 elif "SamplingStrategyType" in default:
81 default = default.split("::")[-1].lower()
82 fields.append((name, default, ctype))
83 except Exception as e:
84 logger.warning(f"解析hpp失败: {e},使用默认字段")
85 return fields
86
87
88def generate_config_json():
89 config = {}
90 comments = {}
91 python_alias = {}
92 for name, default, _, comment, alias in MODEL_CONFIG_FIELDS + CAUSAL_LM_FIELDS:
93 config[name] = default
94 comments[name] = comment
95 if alias:
96 python_alias[name] = alias
97 config["_comment"] = comments
98 config["_python_alias"] = python_alias
99 return config
100
101
102def generate_special_tokens_map():
103 return {
104 "pad_token": "<pad>",
105 "pad_token_id": 0,
106 "bos_token": "<s>",
107 "bos_token_id": 1,
108 "eos_token": "</s>",
109 "eos_token_id": 2,
110 "unk_token": "<unk>",
111 "unk_token_id": 3,
112 "_comment": {
113 "pad_token": "填充token,用于batch对齐",
114 "bos_token": "序列起始token",
115 "eos_token": "序列结束token,生成时遇到此token停止",
116 "unk_token": "未知token,词表外字符映射到此",
117 },
118 }
119
120
121def generate_generation_config_json():
122 config = {}
123 comments = {}
124 for name, default, _, comment, _ in GENERATE_FIELDS:
125 config[name] = default
126 comments[name] = comment
127 config["_comment"] = comments
128 return config
129
130
131def validate_config(config, schema_type):
132 errors = []
133 if schema_type == "config":
134 if config.get("vocab_size", 0) <= 0:
135 errors.append("vocab_size必须>0")
136 if config.get("input_dim", 0) <= 0:
137 errors.append("input_dim必须>0")
138 elif schema_type == "generation_config":
139 t = config.get("temperature", 0)
140 if t <= 0 or t > 2.0:
141 errors.append(f"temperature须在(0.0,2.0],当前={t}")
142 s = config.get("strategy", "")
143 if s not in ("greedy", "top_k", "top_p", "top_k_top_p"):
144 errors.append(f"strategy须为greedy/top_k/top_p/top_k_top_p,当前={s}")
145 return errors
146
147
148def save_all_configs(output_dir, model_hpp=None, generative_hpp=None):
149 os.makedirs(output_dir, exist_ok=True)
150
151 config = generate_config_json()
152 if model_hpp:
153 parsed = parse_config_from_hpp(model_hpp, "Config")
154 if parsed:
155 for name, default, _ in parsed:
156 if name in config and name not in ("_comment", "_python_alias"):
157 config[name] = default
158 errors = validate_config(config, "config")
159 if errors:
160 logger.warning(f"config.json校验问题: {errors}")
161 with open(os.path.join(output_dir, "config.json"), "w", encoding="utf-8") as f:
162 json.dump(config, f, indent=2, ensure_ascii=False)
163
164 special = generate_special_tokens_map()
165 with open(os.path.join(output_dir, "special_tokens_map.json"), "w", encoding="utf-8") as f:
166 json.dump(special, f, indent=2, ensure_ascii=False)
167
168 gen_config = generate_generation_config_json()
169 if generative_hpp:
170 parsed = parse_config_from_hpp(generative_hpp, "GenerateConfig")
171 if parsed:
172 for name, default, _ in parsed:
173 if name in gen_config and name not in ("_comment",):
174 gen_config[name] = default
175 errors = validate_config(gen_config, "generation_config")
176 if errors:
177 logger.warning(f"generation_config.json校验问题: {errors}")
178 with open(os.path.join(output_dir, "generation_config.json"), "w", encoding="utf-8") as f:
179 json.dump(gen_config, f, indent=2, ensure_ascii=False)
180
181 logger.info(f"配置文件已保存到 {output_dir}")
182
183
184if __name__ == "__main__":
185 logging.basicConfig(level=logging.INFO)
186 parser = argparse.ArgumentParser(description="NeuroFlow配置文件生成器")
187 parser.add_argument("--model-hpp", type=str, default="", help="model.hpp路径")
188 parser.add_argument("--generative-hpp", type=str, default="", help="generative.hpp路径")
189 parser.add_argument("--output-dir", type=str, default="configs", help="输出目录")
190 parser.add_argument("--reference-config", type=str, default="", help="Python训练系统参考配置")
191 args = parser.parse_args()
192 save_all_configs(args.output_dir, args.model_hpp or None, args.generative_hpp or None)