用户:
Nomen查看:0 回复:1 评论:0 创建时间:2023-07-26T08:27:57
这里是没有开场白的Nomen,NeuralNetwork的权重保存问题一直是一件很大的问题,由于神经网络本身形状复杂,并且神经突触数量庞大,一个4*128*128*2的神经网络就有17152个神经突触,每一个神经突触都有一个0~16位浮点数作为权重值,那么换算下来权重集的文件将是巨大的,尤其是如果每一个神经突触都附带一串id作为标识的话,那么一个权重集文件就会有差不多1~2mb的大小
但是我们带来了一个全新的,快速的保存与载入的解决方案,平均载入53504个连接只需要150毫秒,如果有条件可以用fs写入一个保存好的文件
// 将储存的字符串载入到实例中
applyParsingConnections("保存的权重集,例如:1.4300415005335936,-1.2477414972574556,0.2487841133735087,-1.7313414634188216,1.7820444680726566...", /* 一个神经网络实例 */);
// 将神经网络权重储存为字符串
fs.writeFile("./src/weights.txt", getConnectionsOut(/* 一个神经网络实例 */))
注:如果要使用该代码,您的NeuralNetwork版本必须大于等于1.1.0,您可以在Box3-Tools找到最新版本的模块
function parseWeightsSet(c, u) {
if (!c.length) return [];
let r = c.split(",");
const net = {
in喵idden"][0].length,
hidden: r.length - u.net["in"].length * u.net["hidden"][0].length - u.net["out"].length * u.net["hidden"][u.net["hidden"].length - 1].length,
out: u.net["out"].length * u.net["hidden"][u.net["hidden"].length - 1].length
};
let ntse = [
r.slice(0, net.in),
r.slice(net.in, net.in + net.hidden),
r.slice(net.in + net.hidden, net.in + net.hidden + net.out)
];
let output = [];
ntse.flat(3).forEach((v, i, a) => {
output.push({
...getLinkInfo(i, [u.net["in"].length, u.net["hidden"].map(v => v.length), u.net["out"].length]),
w: parseFloat(v)
})
})
return output;
}
function getLinkInfo(i, u) {
const input = u[0];
const out = u[2];
const e = u[1];
if (i < input * e[0]) {
let inputNode = Math.floor(i / e[0]);
let hiddenNode = i % e[0];
return {
a: `INPUT-${inputNode}`,
b: `HIDDEN-0-${hiddenNode}`
};
} else {
i -= input * e[0];
for (let layer = 0; layer < e.length - 1; layer++) {
if (i < e[layer] * e[layer + 1]) {
let node1 = Math.floor(i / e[layer + 1]);
let node2 = i % e[layer + 1];
return {
a: `HIDDEN-${layer}-${node1}`,
b: `HIDDEN-${layer + 1}-${node2}`
}
} else {
i -= e[layer] * e[layer + 1];
}
}
if (i < e[e.length - 1] * out) {
let hiddenNode = Math.floor(i / out);
let outputNode = i % out;
return {
a: `HIDDEN-${e.length - 1}-${hiddenNode}`,
b: `OUTPUT-${outputNode}`
}
} else {
throw new Error("i is too large");
}
}
}
function applyParsingConnections(cn, u) {
let r = parseWeightsSet(cn.toString(), u);
u.connections = [];
u.forEachNeuron(n => n.cleanConnections());
let NeuronsMap = {};
u.forEachNeuron(v => NeuronsMap[v.getId()] = v);
r.forEach(v => {
u.connections.push(new Connection(`${v.a}--${v.b}`, NeuronsMap[v.a], NeuronsMap[v.b], v.w));
})
u.connections.forEach(v => v.connect())
}
function getConnectionsOut(u) {
return u.connections.map(v => v.getWeight()).join(",");
}