Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
config_generator.py192 linesDownload Raw Back to scripts
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)