Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
weight_io.cpp361 linesDownload Raw Back to src
1#include "weight_io.hpp"
2#include <iostream>
3#include <sstream>
4#include <cstring>
5#include <algorithm>
6#include <numeric>
7
8namespace neuroflow {
9
10void WeightInitializer::xavier_uniform(Tensor& weight, size_t fan_in, size_t fan_out, std::mt19937& rng) {
11    if (weight.numel() == 0) return;
12    if (fan_in == 0 || fan_out == 0) {
13        std::cerr << "Warning: xavier_uniform with fan_in=" << fan_in << " fan_out=" << fan_out << ", skipping" << std::endl;
14        return;
15    }
16    float a = std::sqrt(6.0f / static_cast<float>(fan_in + fan_out));
17    std::uniform_real_distribution<float> dist(-a, a);
18    float* data = weight.as_fp32();
19    for (size_t i = 0; i < weight.numel(); ++i) {
20        data[i] = dist(rng);
21    }
22}
23
24void WeightInitializer::kaiming_normal(Tensor& weight, size_t fan_in, std::mt19937& rng) {
25    if (weight.numel() == 0) return;
26    if (fan_in == 0) {
27        std::cerr << "Warning: kaiming_normal with fan_in=0, skipping" << std::endl;
28        return;
29    }
30    float std_dev = std::sqrt(2.0f / static_cast<float>(fan_in));
31    std::normal_distribution<float> dist(0.0f, std_dev);
32    float* data = weight.as_fp32();
33    for (size_t i = 0; i < weight.numel(); ++i) {
34        data[i] = dist(rng);
35    }
36}
37
38void WeightInitializer::zeros(Tensor& tensor) {
39    if (tensor.numel() == 0) return;
40    std::memset(tensor.as_fp32(), 0, tensor.data_size_);
41}
42
43void WeightInitializer::random_normal(Tensor& tensor, float mean, float std_dev, std::mt19937& rng) {
44    if (tensor.numel() == 0) return;
45    if (std_dev <= 0.0f) {
46        std::cerr << "Warning: random_normal with std_dev<=0, using 0.01" << std::endl;
47        std_dev = 0.01f;
48    }
49    std::normal_distribution<float> dist(mean, std_dev);
50    float* data = tensor.as_fp32();
51    for (size_t i = 0; i < tensor.numel(); ++i) {
52        data[i] = dist(rng);
53    }
54}
55
56void WeightInitializer::init_model_weights(NeuroFlowModel& model, InitStrategy strategy, uint32_t seed) {
57    std::mt19937 rng(seed);
58
59    auto init_linear = [&](std::shared_ptr<Linear>& layer, size_t fan_in, size_t fan_out) {
60        if (!layer) return;
61        switch (strategy) {
62            case InitStrategy::XAVIER_UNIFORM:
63                xavier_uniform(layer->weight, fan_in, fan_out, rng);
64                break;
65            case InitStrategy::KAIMING_NORMAL:
66                kaiming_normal(layer->weight, fan_in, rng);
67                break;
68            case InitStrategy::ZEROS:
69                zeros(layer->weight);
70                break;
71            case InitStrategy::RANDOM_NORMAL:
72                random_normal(layer->weight, 0.0f, 0.02f, rng);
73                break;
74            default:
75                xavier_uniform(layer->weight, fan_in, fan_out, rng);
76                break;
77        }
78        zeros(layer->bias);
79    };
80
81    auto& cfg = model.config;
82    size_t half = cfg.hidden_dim / 2;
83
84    // === Input Projection ===
85    init_linear(model.input_proj_linear, cfg.input_dim, cfg.hidden_dim);
86
87    // === ECN (Executive Control Network) ===
88    // dlPFC: first layer input_dim=hidden_dim, subsequent=hidden_dim
89    for (size_t i = 0; i < model.ecn->dlpfc_linear.size(); ++i) {
90        size_t fan_in = (i == 0) ? cfg.hidden_dim : cfg.hidden_dim;
91        init_linear(model.ecn->dlpfc_linear[i], fan_in, cfg.hidden_dim);
92    }
93    // OFC: hidden_dim -> half -> 1
94    init_linear(model.ecn->ofc1, cfg.hidden_dim, half);
95    init_linear(model.ecn->ofc2, half, 1);
96    // vmPFC: hidden_dim -> half -> output_dim
97    init_linear(model.ecn->vmpfc1, cfg.hidden_dim, half);
98    init_linear(model.ecn->vmpfc2, half, cfg.hidden_dim);
99
100    // === DMN (Default Mode Network) ===
101    // mem_encoder: memory_dim -> latent_dim*2 -> latent_dim
102    // latent_dim = hidden_dim/2
103    size_t latent_dim = half;
104    init_linear(model.dmn->mem_encoder1, cfg.memory_dim, latent_dim * 2);
105    init_linear(model.dmn->mem_encoder2, latent_dim * 2, latent_dim);
106    // association heads: latent_dim -> latent_dim (each)
107    for (auto& [h1, h2] : model.dmn->association_heads) {
108        init_linear(h1, latent_dim, latent_dim);
109        init_linear(h2, latent_dim, latent_dim);
110    }
111    // future_proj: latent_dim * num_assoc -> latent_dim * 2
112    init_linear(model.dmn->future_proj1, latent_dim * cfg.num_associations, latent_dim * 2);
113
114    // === SN (Salience Network) ===
115    size_t sn_hidden = half;
116    init_linear(model.sn->saliency1, cfg.hidden_dim, sn_hidden);
117    init_linear(model.sn->saliency2, sn_hidden, sn_hidden / 2);
118    init_linear(model.sn->saliency3, sn_hidden / 2, 1);
119    init_linear(model.sn->gate1, cfg.hidden_dim, sn_hidden);
120    init_linear(model.sn->gate2, sn_hidden, 2);
121    init_linear(model.sn->anomaly1, cfg.hidden_dim, sn_hidden);
122    init_linear(model.sn->anomaly2, sn_hidden, 1);
123
124    // === Memory Consolidation Module ===
125    init_linear(model.memory->encode_proj, cfg.hidden_dim, cfg.memory_dim);
126    init_linear(model.memory->retrieve_proj, cfg.memory_dim, cfg.hidden_dim);
127    init_linear(model.memory->query_proj, cfg.hidden_dim, cfg.memory_dim);
128    zeros(model.memory->memory_bank);
129
130    // === Manifold Projection ===
131    size_t manifold_in = cfg.hidden_dim + half;
132    init_linear(model.manifold_proj1, manifold_in, cfg.hidden_dim);
133    init_linear(model.manifold_proj2, cfg.hidden_dim, 32);
134
135    // === Output Fusion (低秩因式分解) ===
136    size_t fusion_in = cfg.hidden_dim * 3;
137    size_t bn = cfg.fusion_bottleneck_dim;
138    init_linear(model.output_fusion_down, fusion_in, bn);
139    init_linear(model.output_fusion_up, bn, cfg.hidden_dim);
140}
141
142ValidationResult WeightInitializer::validate_dimensions(const NeuroFlowModel& model) {
143    ValidationResult result;
144    auto& cfg = model.config;
145    size_t half = cfg.hidden_dim / 2;
146    size_t latent_dim = half;
147
148    auto check2d = [&](const std::string& name, const Tensor& t, size_t expected_rows, size_t expected_cols) {
149        if (t.shape_.size() < 2 || t.shape_[0] != expected_rows || t.shape_[1] != expected_cols) {
150            result.all_passed = false;
151            std::string actual = (t.shape_.size() >= 2)
152                ? "[" + std::to_string(t.shape_[0]) + "," + std::to_string(t.shape_[1]) + "]"
153                : "ndim=" + std::to_string(t.shape_.size());
154            result.failures.push_back(name + ": 期望[" + std::to_string(expected_rows) + "," +
155                std::to_string(expected_cols) + "], 实际" + actual);
156        }
157    };
158
159    auto check1d = [&](const std::string& name, const Tensor& t, size_t expected_dim) {
160        if (t.shape_.size() < 1 || t.shape_[0] != expected_dim) {
161            result.all_passed = false;
162            result.failures.push_back(name + ": 期望[" + std::to_string(expected_dim) + "], 实际dim不匹配");
163        }
164    };
165
166    // Input Projection
167    check2d("input_proj.weight", model.input_proj_linear->weight, cfg.hidden_dim, cfg.input_dim);
168    check1d("input_proj.bias", model.input_proj_linear->bias, cfg.hidden_dim);
169
170    // ECN dlPFC
171    for (size_t i = 0; i < model.ecn->dlpfc_linear.size(); ++i) {
172        check2d("ecn.dlpfc" + std::to_string(i) + ".weight",
173                model.ecn->dlpfc_linear[i]->weight, cfg.hidden_dim, cfg.hidden_dim);
174    }
175    check2d("ecn.ofc1.weight", model.ecn->ofc1->weight, half, cfg.hidden_dim);
176    check2d("ecn.ofc2.weight", model.ecn->ofc2->weight, 1, half);
177    check2d("ecn.vmpfc1.weight", model.ecn->vmpfc1->weight, half, cfg.hidden_dim);
178    check2d("ecn.vmpfc2.weight", model.ecn->vmpfc2->weight, cfg.hidden_dim, half);
179
180    // DMN
181    check2d("dmn.mem_encoder1.weight", model.dmn->mem_encoder1->weight, latent_dim * 2, cfg.memory_dim);
182    check2d("dmn.mem_encoder2.weight", model.dmn->mem_encoder2->weight, latent_dim, latent_dim * 2);
183    for (size_t i = 0; i < model.dmn->association_heads.size(); ++i) {
184        auto& [h1, h2] = model.dmn->association_heads[i];
185        check2d("dmn.head" + std::to_string(i) + ".1.weight", h1->weight, latent_dim, latent_dim);
186        check2d("dmn.head" + std::to_string(i) + ".2.weight", h2->weight, latent_dim, latent_dim);
187    }
188    check2d("dmn.future_proj1.weight", model.dmn->future_proj1->weight, latent_dim * 2, latent_dim * cfg.num_associations);
189
190    // SN
191    size_t sn_hidden = half;
192    check2d("sn.saliency1.weight", model.sn->saliency1->weight, sn_hidden, cfg.hidden_dim);
193    check2d("sn.saliency2.weight", model.sn->saliency2->weight, sn_hidden / 2, sn_hidden);
194    check2d("sn.saliency3.weight", model.sn->saliency3->weight, 1, sn_hidden / 2);
195    check2d("sn.gate1.weight", model.sn->gate1->weight, sn_hidden, cfg.hidden_dim);
196    check2d("sn.gate2.weight", model.sn->gate2->weight, 2, sn_hidden);
197    check2d("sn.anomaly1.weight", model.sn->anomaly1->weight, sn_hidden, cfg.hidden_dim);
198    check2d("sn.anomaly2.weight", model.sn->anomaly2->weight, 1, sn_hidden);
199
200    // Memory
201    check2d("memory.encode_proj.weight", model.memory->encode_proj->weight, cfg.memory_dim, cfg.hidden_dim);
202    check2d("memory.retrieve_proj.weight", model.memory->retrieve_proj->weight, cfg.hidden_dim, cfg.memory_dim);
203    check2d("memory.query_proj.weight", model.memory->query_proj->weight, cfg.memory_dim, cfg.hidden_dim);
204
205    // Manifold
206    size_t manifold_in = cfg.hidden_dim + half;
207    check2d("manifold_proj1.weight", model.manifold_proj1->weight, cfg.hidden_dim, manifold_in);
208    check2d("manifold_proj2.weight", model.manifold_proj2->weight, 32, cfg.hidden_dim);
209
210    // Output Fusion (低秩因式分解)
211    size_t bn = cfg.fusion_bottleneck_dim;
212    check2d("output_fusion.down.weight", model.output_fusion_down->weight, bn, cfg.hidden_dim * 3);
213    check1d("output_fusion.down.bias", model.output_fusion_down->bias, bn);
214    check2d("output_fusion.up.weight", model.output_fusion_up->weight, cfg.hidden_dim, bn);
215    check1d("output_fusion.up.bias", model.output_fusion_up->bias, cfg.hidden_dim);
216
217    return result;
218}
219
220void save_binary(const NeuroFlowModel& model, const std::string& path) {
221    model.save(path);
222}
223
224void load_binary(NeuroFlowModel& model, const std::string& path) {
225    model.load(path);
226}
227
228void save_npz(const NeuroFlowModel& model, const std::string& path) {
229    // NPZ格式需要zlib/miniz依赖
230    // 当前实现:将每个权重层保存为独立的.raw文件 + 一个manifest.json索引
231    // 这与NumPy的.npz格式(ZIP存档)兼容性有限
232    // 完整NPZ实现需引入cnpy库: https://github.com/rogersce/cnpy
233    //
234    // 回退策略:使用NFv1二进制格式
235    model.save(path + ".nfv1");
236    std::cerr << "注意: NPZ格式保存需要cnpy/zlib依赖,当前回退到NFv1格式" << std::endl;
237    std::cerr << "  如需完整NPZ支持,请安装cnpy: https://github.com/rogersce/cnpy" << std::endl;
238}
239
240void load_npz(NeuroFlowModel& model, const std::string& path) {
241    // 同save_npz,回退到NFv1
242    model.load(path + ".nfv1");
243}
244
245void save_metadata(const NeuroFlowModel& model, const std::string& path) {
246    std::ofstream ofs(path);
247    if (!ofs) {
248        std::cerr << "Warning: cannot open metadata file: " << path << std::endl;
249        return;
250    }
251
252    auto& cfg = model.config;
253    size_t half = cfg.hidden_dim / 2;
254    size_t latent_dim = half;
255
256    auto count_params = [](const Tensor& t) -> size_t { return t.numel(); };
257
258    ofs << "{\n";
259    ofs << "  \"model_type\": \"NeuroFlow\",\n";
260    ofs << "  \"format\": \"NFv1\",\n";
261    ofs << "  \"config\": {\n";
262    ofs << "    \"input_dim\": " << cfg.input_dim << ",\n";
263    ofs << "    \"hidden_dim\": " << cfg.hidden_dim << ",\n";
264    ofs << "    \"output_dim\": " << cfg.output_dim << ",\n";
265    ofs << "    \"memory_dim\": " << cfg.memory_dim << ",\n";
266    ofs << "    \"memory_slots\": " << cfg.memory_slots << ",\n";
267    ofs << "    \"num_layers\": " << cfg.num_layers << ",\n";
268    ofs << "    \"num_associations\": " << cfg.num_associations << ",\n";
269    ofs << "    \"vocab_size\": " << cfg.vocab_size << ",\n";
270    ofs << "    \"max_seq_len\": " << cfg.max_seq_len << "\n";
271    ofs << "  },\n";
272
273    ofs << "  \"layers\": [\n";
274    auto add_layer = [&](const std::string& name, const std::string& type,
275                         const Tensor& weight, const Tensor& bias, bool& first) {
276        if (!first) ofs << ",\n";
277        first = false;
278        ofs << "    {\"name\": \"" << name << "\", \"type\": \"" << type << "\", ";
279        ofs << "\"weight_shape\": [";
280        for (size_t i = 0; i < weight.shape_.size(); ++i) {
281            if (i > 0) ofs << ", ";
282            ofs << weight.shape_[i];
283        }
284        ofs << "], \"params\": " << count_params(weight) + (bias.data_ ? count_params(bias) : 0) << "}";
285    };
286
287    bool first = true;
288    add_layer("input_proj", "Linear", model.input_proj_linear->weight, model.input_proj_linear->bias, first);
289    for (size_t i = 0; i < model.ecn->dlpfc_linear.size(); ++i) {
290        add_layer("ecn.dlpfc" + std::to_string(i), "Linear",
291                  model.ecn->dlpfc_linear[i]->weight, model.ecn->dlpfc_linear[i]->bias, first);
292    }
293    add_layer("ecn.ofc1", "Linear", model.ecn->ofc1->weight, model.ecn->ofc1->bias, first);
294    add_layer("ecn.ofc2", "Linear", model.ecn->ofc2->weight, model.ecn->ofc2->bias, first);
295    add_layer("ecn.vmpfc1", "Linear", model.ecn->vmpfc1->weight, model.ecn->vmpfc1->bias, first);
296    add_layer("ecn.vmpfc2", "Linear", model.ecn->vmpfc2->weight, model.ecn->vmpfc2->bias, first);
297    add_layer("dmn.mem_encoder1", "Linear", model.dmn->mem_encoder1->weight, model.dmn->mem_encoder1->bias, first);
298    add_layer("dmn.mem_encoder2", "Linear", model.dmn->mem_encoder2->weight, model.dmn->mem_encoder2->bias, first);
299    for (size_t i = 0; i < model.dmn->association_heads.size(); ++i) {
300        auto& [h1, h2] = model.dmn->association_heads[i];
301        add_layer("dmn.head" + std::to_string(i) + ".1", "Linear", h1->weight, h1->bias, first);
302        add_layer("dmn.head" + std::to_string(i) + ".2", "Linear", h2->weight, h2->bias, first);
303    }
304    add_layer("dmn.future_proj1", "Linear", model.dmn->future_proj1->weight, model.dmn->future_proj1->bias, first);
305    add_layer("sn.saliency1", "Linear", model.sn->saliency1->weight, model.sn->saliency1->bias, first);
306    add_layer("sn.saliency2", "Linear", model.sn->saliency2->weight, model.sn->saliency2->bias, first);
307    add_layer("sn.saliency3", "Linear", model.sn->saliency3->weight, model.sn->saliency3->bias, first);
308    add_layer("sn.gate1", "Linear", model.sn->gate1->weight, model.sn->gate1->bias, first);
309    add_layer("sn.gate2", "Linear", model.sn->gate2->weight, model.sn->gate2->bias, first);
310    add_layer("sn.anomaly1", "Linear", model.sn->anomaly1->weight, model.sn->anomaly1->bias, first);
311    add_layer("sn.anomaly2", "Linear", model.sn->anomaly2->weight, model.sn->anomaly2->bias, first);
312    add_layer("memory.encode_proj", "Linear", model.memory->encode_proj->weight, model.memory->encode_proj->bias, first);
313    add_layer("memory.retrieve_proj", "Linear", model.memory->retrieve_proj->weight, model.memory->retrieve_proj->bias, first);
314    add_layer("memory.query_proj", "Linear", model.memory->query_proj->weight, model.memory->query_proj->bias, first);
315    add_layer("manifold_proj1", "Linear", model.manifold_proj1->weight, model.manifold_proj1->bias, first);
316    add_layer("manifold_proj2", "Linear", model.manifold_proj2->weight, model.manifold_proj2->bias, first);
317    add_layer("output_fusion.down", "Linear", model.output_fusion_down->weight, model.output_fusion_down->bias, first);
318    add_layer("output_fusion.up", "Linear", model.output_fusion_up->weight, model.output_fusion_up->bias, first);
319    ofs << "\n  ],\n";
320
321    size_t total_params = 0;
322    auto count_layer = [&](const Tensor& w, const Tensor& b) {
323        total_params += count_params(w) + (b.data_ ? count_params(b) : 0);
324    };
325    count_layer(model.input_proj_linear->weight, model.input_proj_linear->bias);
326    for (auto& l : model.ecn->dlpfc_linear) count_layer(l->weight, l->bias);
327    count_layer(model.ecn->ofc1->weight, model.ecn->ofc1->bias);
328    count_layer(model.ecn->ofc2->weight, model.ecn->ofc2->bias);
329    count_layer(model.ecn->vmpfc1->weight, model.ecn->vmpfc1->bias);
330    count_layer(model.ecn->vmpfc2->weight, model.ecn->vmpfc2->bias);
331    count_layer(model.dmn->mem_encoder1->weight, model.dmn->mem_encoder1->bias);
332    count_layer(model.dmn->mem_encoder2->weight, model.dmn->mem_encoder2->bias);
333    for (auto& [h1, h2] : model.dmn->association_heads) {
334        count_layer(h1->weight, h1->bias);
335        count_layer(h2->weight, h2->bias);
336    }
337    count_layer(model.dmn->future_proj1->weight, model.dmn->future_proj1->bias);
338    count_layer(model.sn->saliency1->weight, model.sn->saliency1->bias);
339    count_layer(model.sn->saliency2->weight, model.sn->saliency2->bias);
340    count_layer(model.sn->saliency3->weight, model.sn->saliency3->bias);
341    count_layer(model.sn->gate1->weight, model.sn->gate1->bias);
342    count_layer(model.sn->gate2->weight, model.sn->gate2->bias);
343    count_layer(model.sn->anomaly1->weight, model.sn->anomaly1->bias);
344    count_layer(model.sn->anomaly2->weight, model.sn->anomaly2->bias);
345    count_layer(model.memory->encode_proj->weight, model.memory->encode_proj->bias);
346    count_layer(model.memory->retrieve_proj->weight, model.memory->retrieve_proj->bias);
347    count_layer(model.memory->query_proj->weight, model.memory->query_proj->bias);
348    count_layer(model.manifold_proj1->weight, model.manifold_proj1->bias);
349    count_layer(model.manifold_proj2->weight, model.manifold_proj2->bias);
350    count_layer(model.output_fusion_down->weight, model.output_fusion_down->bias);
351    count_layer(model.output_fusion_up->weight, model.output_fusion_up->bias);
352
353    ofs << "  \"total_params\": " << total_params << ",\n";
354    ofs << "  \"memory_bank_slots\": " << cfg.memory_slots << ",\n";
355    ofs << "  \"memory_bank_dim\": " << cfg.memory_dim << "\n";
356    ofs << "}\n";
357    ofs.close();
358}
359
360}
361