bit-jev 0.8.7__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
bit_jev/__init__.py ADDED
@@ -0,0 +1,18 @@
1
+ """BitNet 结构化决策:训练、I2_S 导出与逐题 CPU 推理共用同一编码契约。"""
2
+
3
+ # 包版本与 pyproject.toml、项目更新日志保持同步。
4
+ __version__ = "0.8.7"
5
+
6
+ # 默认骨干用于新训练运行;蒸馏和 CPU 推理从检查点读取实际来源。
7
+ DEFAULT_BASE = "microsoft/bitnet-b1.58-2B-4T-bf16"
8
+
9
+ # 顶层 API 延迟导入,单纯查询包版本时不会初始化 tokenizer 或原生构建器。
10
+ __all__ = ["BitJev", "download_model", "__version__"]
11
+
12
+
13
+ def __getattr__(name):
14
+ """按需公开 GGUF 常驻模型与模型下载入口。"""
15
+ if name in {"BitJev", "download_model"}:
16
+ from .gguf import BitJev, download_model
17
+ return {"BitJev": BitJev, "download_model": download_model}[name]
18
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
bit_jev/__version__.py ADDED
@@ -0,0 +1,2 @@
1
+ # 源码包版本与 pyproject.toml 中的发行版本保持一致。
2
+ __version__ = "0.8.7"
@@ -0,0 +1,26 @@
1
+ cmake_minimum_required(VERSION 3.28)
2
+ project(bit_jev_native LANGUAGES C CXX)
3
+
4
+ # 只编译本项目所需的 GGUF 推理库与命令行程序;后端由调用方选择。
5
+ set(LLAMA_BUILD_TOOLS OFF CACHE BOOL "" FORCE)
6
+ set(LLAMA_BUILD_EXAMPLES OFF CACHE BOOL "" FORCE)
7
+ set(LLAMA_BUILD_TESTS OFF CACHE BOOL "" FORCE)
8
+ set(BUILD_SHARED_LIBS OFF CACHE BOOL "" FORCE)
9
+
10
+ # 源码构建默认读取工作区依赖;pip 缓存构建显式传入固定版本的路径。
11
+ if(NOT DEFINED BITNET_CPP_DIR)
12
+ set(BITNET_CPP_DIR "${CMAKE_CURRENT_LIST_DIR}/../../learning/bitnet/3rdparty/llama.cpp")
13
+ endif()
14
+ add_subdirectory("${BITNET_CPP_DIR}" "${CMAKE_BINARY_DIR}/bitnet_cpp")
15
+
16
+ # 原生程序读取已编码题目与 float32 指针头,主干推理由 llama 库承担。
17
+ add_executable(bit-jev-cpu main.cpp)
18
+ target_compile_features(bit-jev-cpu PRIVATE cxx_std_17)
19
+ # 设备请求必须与编译出的 GPU 后端一致,避免外部 binary 参数静默换后端。
20
+ if(GGML_CUDA)
21
+ target_compile_definitions(bit-jev-cpu PRIVATE BIT_JEV_BUILT_CUDA=1)
22
+ elseif(GGML_VULKAN)
23
+ target_compile_definitions(bit-jev-cpu PRIVATE BIT_JEV_BUILT_VULKAN=1)
24
+ endif()
25
+ target_include_directories(bit-jev-cpu PRIVATE "${BITNET_CPP_DIR}/include" "${BITNET_CPP_DIR}/vendor")
26
+ target_link_libraries(bit-jev-cpu PRIVATE llama)
@@ -0,0 +1,13 @@
1
+ diff --git a/src/models/bitnet.cpp b/src/models/bitnet.cpp
2
+ --- a/src/models/bitnet.cpp
3
+ +++ b/src/models/bitnet.cpp
4
+ @@ -130,7 +130,8 @@ llama_model_bitnet::graph::graph(const llama_model & model, const llm_graph_para
5
+ model.layers[il].ffn_gate, NULL, model.layers[il].ffn_gate_s,
6
+ NULL, NULL, NULL,
7
+ NULL,
8
+ - LLM_FFN_SILU, LLM_FFN_PAR, il);
9
+ + // 蒸馏检查点的 hidden_act 为 relu2,必须与 Transformers 前向保持一致。
10
+ + LLM_FFN_RELU_SQR, LLM_FFN_PAR, il);
11
+ cb(cur, "ffn_sub_out", il);
12
+
13
+ cur = build_norm(cur,
@@ -0,0 +1,283 @@
1
+ // bit-jev 原生 GGUF 推理:逐题运行 BitNet,并在候选边界读取隐藏状态。
2
+ #include "llama.h"
3
+ #include "ggml-backend.h"
4
+ #include "nlohmann/json.hpp"
5
+
6
+ #include <algorithm>
7
+ #include <chrono>
8
+ #include <cmath>
9
+ #include <cstdint>
10
+ #include <cstring>
11
+ #include <fstream>
12
+ #include <iostream>
13
+ #include <memory>
14
+ #include <stdexcept>
15
+ #include <string>
16
+ #include <unordered_set>
17
+ #include <vector>
18
+
19
+ using Json = nlohmann::json;
20
+
21
+ // 模型加载时仅输出警告和错误,避免每次启动打印整份张量目录。
22
+ void native_log(ggml_log_level level, const char * message, void *) {
23
+ if (level >= GGML_LOG_LEVEL_WARN) std::cerr << message;
24
+ }
25
+
26
+ // 指针头文件的矩阵均按行优先排列,顺序与 Python 导出器一致。
27
+ struct Head {
28
+ uint32_t hidden_size = 0;
29
+ uint32_t head_size = 0;
30
+ float temperature = 1.0f;
31
+ std::vector<float> q_weight;
32
+ std::vector<float> q_bias;
33
+ std::vector<float> k_weight;
34
+ std::vector<float> k_bias;
35
+ };
36
+
37
+ // llama_batch 持有 C API 分配的内存,在异常路径也自动释放。
38
+ struct BatchHolder {
39
+ explicit BatchHolder(int capacity) : batch(llama_batch_init(capacity, 0, 1)) {}
40
+ ~BatchHolder() { llama_batch_free(batch); }
41
+ llama_batch batch;
42
+ };
43
+
44
+ // 从固定布局的 sidecar 读取全部指针头参数。
45
+ Head load_head(const std::string & path) {
46
+ std::ifstream source(path, std::ios::binary);
47
+ if (!source) throw std::runtime_error("无法打开 head.f32:" + path);
48
+ char magic[8] = {};
49
+ Head head;
50
+ source.read(magic, 8);
51
+ source.read(reinterpret_cast<char *>(&head.hidden_size), sizeof(head.hidden_size));
52
+ source.read(reinterpret_cast<char *>(&head.head_size), sizeof(head.head_size));
53
+ source.read(reinterpret_cast<char *>(&head.temperature), sizeof(head.temperature));
54
+ if (!source || std::memcmp(magic, "BJHEAD01", 8) != 0 || head.hidden_size == 0 ||
55
+ head.head_size == 0 || !std::isfinite(head.temperature) || head.temperature <= 0) {
56
+ throw std::runtime_error("head.f32 文件头无效");
57
+ }
58
+ // 四段数组的尺寸由文件头决定,读取长度不足即视为损坏。
59
+ const size_t matrix_size = static_cast<size_t>(head.hidden_size) * head.head_size;
60
+ auto read_array = [&source](size_t count) {
61
+ std::vector<float> values(count);
62
+ source.read(reinterpret_cast<char *>(values.data()), count * sizeof(float));
63
+ if (!source) throw std::runtime_error("head.f32 参数数据截断");
64
+ return values;
65
+ };
66
+ head.q_weight = read_array(matrix_size);
67
+ head.q_bias = read_array(head.head_size);
68
+ head.k_weight = read_array(matrix_size);
69
+ head.k_bias = read_array(head.head_size);
70
+ if (source.peek() != std::char_traits<char>::eof()) {
71
+ throw std::runtime_error("head.f32 含有多余数据");
72
+ }
73
+ return head;
74
+ }
75
+
76
+ // 将一个骨干隐藏状态投影到指针空间。
77
+ std::vector<float> project(const std::vector<float> & hidden,
78
+ const std::vector<float> & weights,
79
+ const std::vector<float> & bias,
80
+ const Head & head) {
81
+ std::vector<float> result(head.head_size);
82
+ for (size_t row = 0; row < head.head_size; ++row) {
83
+ float value = bias[row];
84
+ const size_t offset = row * head.hidden_size;
85
+ for (size_t col = 0; col < head.hidden_size; ++col) {
86
+ value += weights[offset + col] * hidden[col];
87
+ }
88
+ result[row] = value;
89
+ }
90
+ return result;
91
+ }
92
+
93
+ // 清空上一题的 KV 缓存后,按普通因果顺序运行一行。
94
+ std::vector<std::vector<float>> read_hidden(llama_context * context, const Json & row,
95
+ uint32_t hidden_size, int batch_size) {
96
+ const std::vector<llama_token> ids = row.at("ids").get<std::vector<llama_token>>();
97
+ const std::vector<int> option_positions = row.at("options").get<std::vector<int>>();
98
+ const int decide_position = row.at("decide").get<int>();
99
+ if (ids.empty() || option_positions.empty() || decide_position != static_cast<int>(ids.size()) - 1 ||
100
+ ids.size() > 4096) {
101
+ throw std::runtime_error("题目 token 序列或 decide 位置无效");
102
+ }
103
+ // 每个输出位置都应位于当前因果行内且互不重复。
104
+ std::unordered_set<int> wanted(option_positions.begin(), option_positions.end());
105
+ wanted.insert(decide_position);
106
+ if (wanted.size() != option_positions.size() + 1) {
107
+ throw std::runtime_error("候选结束位置重复或与 decide 重叠");
108
+ }
109
+ for (int position : wanted) {
110
+ if (position < 0 || position >= static_cast<int>(ids.size())) {
111
+ throw std::runtime_error("隐藏状态读取位置越界");
112
+ }
113
+ }
114
+ llama_memory_clear(llama_get_memory(context), true);
115
+ BatchHolder holder(batch_size);
116
+ std::vector<std::vector<float>> hidden(ids.size());
117
+ // 分块送入模型;每次 decode 后立即复制本块要求的隐藏状态。
118
+ for (size_t start = 0; start < ids.size(); start += batch_size) {
119
+ const int count = static_cast<int>(std::min<size_t>(batch_size, ids.size() - start));
120
+ holder.batch.n_tokens = count;
121
+ for (int index = 0; index < count; ++index) {
122
+ const int position = static_cast<int>(start) + index;
123
+ holder.batch.token[index] = ids[position];
124
+ holder.batch.pos[index] = position;
125
+ holder.batch.n_seq_id[index] = 1;
126
+ holder.batch.seq_id[index][0] = 0;
127
+ // embeddings 模式要求当前块每个 token 都为输出;后续仅复制需要的位置。
128
+ holder.batch.logits[index] = 1;
129
+ }
130
+ if (llama_decode(context, holder.batch) != 0) {
131
+ throw std::runtime_error("BitNet GGUF decode 失败");
132
+ }
133
+ for (int index = 0; index < count; ++index) {
134
+ const int position = static_cast<int>(start) + index;
135
+ if (!wanted.count(position)) continue;
136
+ const float * embedding = llama_get_embeddings_ith(context, index);
137
+ if (embedding == nullptr) throw std::runtime_error("BitNet 未返回请求的隐藏状态");
138
+ hidden[position].assign(embedding, embedding + hidden_size);
139
+ }
140
+ }
141
+ // 结果按候选顺序排列,最后一项是 decide 状态。
142
+ std::vector<std::vector<float>> selected;
143
+ selected.reserve(option_positions.size() + 1);
144
+ for (int position : option_positions) selected.push_back(std::move(hidden[position]));
145
+ selected.push_back(std::move(hidden[decide_position]));
146
+ return selected;
147
+ }
148
+
149
+ // 对候选状态执行与 Python PointerHead 完全相同的投影和温度缩放。
150
+ std::vector<float> score_row(const std::vector<std::vector<float>> & states, const Head & head) {
151
+ const std::vector<float> query = project(states.back(), head.q_weight, head.q_bias, head);
152
+ const float scale = 1.0f / std::sqrt(static_cast<float>(head.head_size)) / head.temperature;
153
+ std::vector<float> logits;
154
+ logits.reserve(states.size() - 1);
155
+ for (size_t option = 0; option + 1 < states.size(); ++option) {
156
+ const std::vector<float> key = project(states[option], head.k_weight, head.k_bias, head);
157
+ float dot = 0.0f;
158
+ for (size_t index = 0; index < head.head_size; ++index) dot += query[index] * key[index];
159
+ logits.push_back(dot * scale);
160
+ }
161
+ return logits;
162
+ }
163
+
164
+ // 使用稳定的 softmax 将一题的 logits 转换为概率。
165
+ std::vector<float> softmax(const std::vector<float> & logits) {
166
+ const float maximum = *std::max_element(logits.begin(), logits.end());
167
+ std::vector<float> probabilities;
168
+ probabilities.reserve(logits.size());
169
+ float sum = 0.0f;
170
+ for (float value : logits) {
171
+ const float probability = std::exp(value - maximum);
172
+ probabilities.push_back(probability);
173
+ sum += probability;
174
+ }
175
+ for (float & probability : probabilities) probability /= sum;
176
+ return probabilities;
177
+ }
178
+
179
+ // 读取命令行选项;所有输入均为显式路径或正整数。
180
+ std::string option_value(int argc, char ** argv, const std::string & name,
181
+ const std::string & fallback = "") {
182
+ for (int index = 1; index + 1 < argc; ++index) {
183
+ if (argv[index] == name) return argv[index + 1];
184
+ }
185
+ return fallback;
186
+ }
187
+
188
+ // 仅检查布尔开关是否出现,不读取其后的路径或数字参数。
189
+ bool has_option(int argc, char ** argv, const std::string & name) {
190
+ for (int index = 1; index < argc; ++index) {
191
+ if (argv[index] == name) return true;
192
+ }
193
+ return false;
194
+ }
195
+
196
+ int main(int argc, char ** argv) {
197
+ const std::string model_path = option_value(argc, argv, "--model");
198
+ const std::string head_path = option_value(argc, argv, "--head");
199
+ const std::string device = option_value(argc, argv, "--device", "cpu");
200
+ const int threads = std::stoi(option_value(argc, argv, "--threads", "4"));
201
+ const int batch_size = std::stoi(option_value(argc, argv, "--batch", "256"));
202
+ if (model_path.empty() || head_path.empty() || threads <= 0 || batch_size <= 0 ||
203
+ (device != "cpu" && device != "vulkan" && device != "cuda")) {
204
+ std::cerr << "用法:bit-jev-cpu --model 文件.gguf --head head.f32 --device cpu|vulkan|cuda --threads 4 --batch 256\n";
205
+ return 2;
206
+ }
207
+ // 外部传入二进制时也强制匹配后端,不能把 Vulkan 程序当 CUDA 程序运行。
208
+ #if !defined(BIT_JEV_BUILT_VULKAN)
209
+ if (device == "vulkan") {
210
+ std::cerr << "本程序未编入 Vulkan 后端\n";
211
+ return 2;
212
+ }
213
+ #endif
214
+ #if !defined(BIT_JEV_BUILT_CUDA)
215
+ if (device == "cuda") {
216
+ std::cerr << "本程序未编入 CUDA 后端\n";
217
+ return 2;
218
+ }
219
+ #endif
220
+ llama_log_set(native_log, nullptr);
221
+ llama_backend_init();
222
+ try {
223
+ // CPU 不卸载层;GPU 模式申请全部层,实际后端由对应的编译选项提供。
224
+ llama_model_params model_params = llama_model_default_params();
225
+ model_params.n_gpu_layers = device == "cpu" ? 0 : 999;
226
+ if (device != "cpu") {
227
+ // 同时检查后端能力和可见 GPU,避免用户要求 GPU 时悄悄在 CPU 上运行。
228
+ if (!llama_supports_gpu_offload() ||
229
+ ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_GPU) == nullptr) {
230
+ throw std::runtime_error("当前原生程序没有可用 GPU 后端或可见 GPU;请检查构建选项和驱动");
231
+ }
232
+ }
233
+ std::unique_ptr<llama_model, decltype(&llama_model_free)> model(
234
+ llama_model_load_from_file(model_path.c_str(), model_params), &llama_model_free);
235
+ if (!model) throw std::runtime_error("无法加载 I2_S GGUF 模型");
236
+ const Head head = load_head(head_path);
237
+ if (head.hidden_size != static_cast<uint32_t>(llama_model_n_embd(model.get()))) {
238
+ throw std::runtime_error("指针头与 BitNet 隐藏维度不一致");
239
+ }
240
+ llama_context_params context_params = llama_context_default_params();
241
+ context_params.n_ctx = 4096;
242
+ context_params.n_batch = batch_size;
243
+ context_params.n_ubatch = batch_size;
244
+ context_params.n_outputs_max = batch_size;
245
+ context_params.n_threads = threads;
246
+ context_params.n_threads_batch = threads;
247
+ context_params.embeddings = true;
248
+ context_params.pooling_type = LLAMA_POOLING_TYPE_NONE;
249
+ std::unique_ptr<llama_context, decltype(&llama_free)> context(
250
+ llama_init_from_model(model.get(), context_params), &llama_free);
251
+ if (!context) throw std::runtime_error("无法初始化 BitNet GGUF 上下文");
252
+ llama_set_embeddings(context.get(), true);
253
+ // Python API 用握手确认模型已经进入内存或显存,再开始计时和请求。
254
+ if (has_option(argc, argv, "--ready")) {
255
+ std::cout << "{\"ready\":true}\n" << std::flush;
256
+ }
257
+ // 模型常驻内存;stdin 每一行独立对应 stdout 的一行结果。
258
+ std::string line;
259
+ while (std::getline(std::cin, line)) {
260
+ if (line.empty()) continue;
261
+ const auto started = std::chrono::steady_clock::now();
262
+ const Json request = Json::parse(line);
263
+ Json result;
264
+ result["logits"] = Json::array();
265
+ result["probabilities"] = Json::array();
266
+ for (const Json & row : request.at("rows")) {
267
+ auto states = read_hidden(context.get(), row, head.hidden_size, batch_size);
268
+ auto logits = score_row(states, head);
269
+ result["logits"].push_back(logits);
270
+ result["probabilities"].push_back(softmax(logits));
271
+ }
272
+ const auto ended = std::chrono::steady_clock::now();
273
+ result["latency_ms"] = std::chrono::duration<double, std::milli>(ended - started).count();
274
+ std::cout << result.dump() << '\n' << std::flush;
275
+ }
276
+ } catch (const std::exception & error) {
277
+ std::cerr << "bit-jev GGUF 推理失败:" << error.what() << '\n';
278
+ llama_backend_free();
279
+ return 1;
280
+ }
281
+ llama_backend_free();
282
+ return 0;
283
+ }
bit_jev/api.py ADDED
@@ -0,0 +1,175 @@
1
+ """TypeSafe-compatible request/response shapes (same contract as kev.api).
2
+
3
+ Noul -> 2 options [false, true]; answer = p(true)
4
+ Choice -> options 'name' or 'name: desc'; answer = argmax + probabilities + confidence
5
+ Score -> options = ordered level descriptions; answer = expected level + probabilities + confidence
6
+ """
7
+ import json
8
+ from typing import Any, Literal, Union
9
+
10
+ JSONContent = Union[str, dict, list, int, float, bool, None]
11
+ MAX_OPTIONS = 255
12
+ QUESTION_TYPES = ("noul", "choice", "score")
13
+
14
+
15
+ def render(v: JSONContent, indent: int = 0) -> str:
16
+ """Flatten str | object | array into text the model sees. Field names are kept as labels."""
17
+ pad = " " * indent
18
+ if v is None:
19
+ return ""
20
+ if isinstance(v, (str, int, float, bool)):
21
+ return str(v)
22
+ if isinstance(v, list):
23
+ return "\n".join(f"{pad}- {render(x, indent + 1).lstrip()}" for x in v)
24
+ return "\n".join(
25
+ f"{pad}{k}:\n{render(x, indent + 1)}" if isinstance(x, (dict, list)) else f"{pad}{k}: {render(x)}"
26
+ for k, x in v.items()
27
+ )
28
+
29
+
30
+ def option_text(name: str, desc: JSONContent) -> str:
31
+ return name if desc is None or desc == "" else f"{name}: {render(desc)}"
32
+
33
+
34
+ def question_keys(qtype: str, criteria) -> list:
35
+ """The keys a question's probabilities are reported under, in option order."""
36
+ if qtype == "choice":
37
+ return list(criteria)
38
+ if qtype == "noul":
39
+ return ["false", "true"]
40
+ return [str(i) for i in range(len(criteria))]
41
+
42
+
43
+ def validate_request(req: dict) -> None:
44
+ """Raise ValueError on a malformed SystemOne-shaped request dict."""
45
+ if not isinstance(req, dict) or "questions" not in req:
46
+ raise ValueError("request needs a 'questions' object")
47
+ if not req["questions"]:
48
+ raise ValueError("at least one question is required")
49
+ for qid, q in req["questions"].items():
50
+ t = q.get("type")
51
+ if t not in QUESTION_TYPES:
52
+ raise ValueError(f"question {qid!r}: type must be one of {QUESTION_TYPES}")
53
+ if t in ("choice", "score"):
54
+ crit = q.get("criteria")
55
+ n = len(crit) if crit else 0
56
+ if not 1 <= n <= MAX_OPTIONS:
57
+ raise ValueError(f"question {qid!r}: criteria must have 1..{MAX_OPTIONS} options")
58
+
59
+
60
+ def to_record(req: dict, labelled: bool = True):
61
+ """SystemOne-shaped dict -> (internal record, per-question meta).
62
+
63
+ With labelled=False, labels are zero placeholders and may be absent."""
64
+ validate_request(req)
65
+ qs, meta = [], []
66
+ for qid, q in req["questions"].items():
67
+ t = q["type"]
68
+ m = {"id": qid, "type": t, "keys": question_keys(t, q.get("criteria"))}
69
+ if t == "noul":
70
+ c = q.get("criteria") or {}
71
+ opts = [option_text("no", c.get("false")), option_text("yes", c.get("true"))]
72
+ label = int(q["label"]) if "label" in q else -1
73
+ elif t == "choice":
74
+ opts = [option_text(k, v) for k, v in q["criteria"].items()]
75
+ label = m["keys"].index(q["label"]) if "label" in q else -1
76
+ else:
77
+ opts = [render(x) for x in q["criteria"]]
78
+ m["legend"] = dict(zip(m["keys"], opts))
79
+ label = int(q["label"]) if "label" in q else -1
80
+ if labelled and label < 0:
81
+ raise ValueError(f"question {qid!r} of type {t} is missing a label")
82
+ qs.append({"instr": render(q.get("instructions")), "options": opts,
83
+ "label": max(label, 0), "qtype": t})
84
+ meta.append(m)
85
+ return {"state": render(req.get("state")), "questions": qs}, meta
86
+
87
+
88
+ def load_requests(path):
89
+ """Labelled requests from a JSONL file: one SystemOne-shaped object per line, a label per question."""
90
+ reqs = []
91
+ with open(path, encoding="utf-8") as f:
92
+ for i, line in enumerate(f, 1):
93
+ line = line.strip()
94
+ if not line:
95
+ continue
96
+ try:
97
+ r = json.loads(line)
98
+ except json.JSONDecodeError as e:
99
+ raise ValueError(f"{path}:{i}: invalid JSON: {e}") from e
100
+ r.setdefault("_meta", {})
101
+ r["_meta"].setdefault("id", f"{path}:{i}")
102
+ r["_meta"].setdefault("source", "user")
103
+ reqs.append(r)
104
+ if not reqs:
105
+ raise ValueError(f"{path}: no records found")
106
+ return reqs
107
+
108
+
109
+ def read_jsonl(path):
110
+ with open(path, encoding="utf-8") as f:
111
+ return [json.loads(line) for line in f if line.strip()]
112
+
113
+
114
+ def write_jsonl(path, rows):
115
+ with open(path, "w", encoding="utf-8", newline="\n") as f:
116
+ for r in rows:
117
+ f.write(json.dumps(r, ensure_ascii=False) + "\n")
118
+
119
+
120
+ def write_json(path, obj):
121
+ with open(path, "w", encoding="utf-8", newline="\n") as f:
122
+ json.dump(obj, f, ensure_ascii=False, indent=2)
123
+
124
+
125
+ def read_json(path):
126
+ with open(path, encoding="utf-8") as f:
127
+ return json.load(f)
128
+
129
+
130
+ def choice_confidence(p) -> float:
131
+ K = len(p)
132
+ return 1.0 if K == 1 else (max(p) - 1 / K) / (1 - 1 / K)
133
+
134
+
135
+ def score_confidence(p) -> float:
136
+ """1 - E|level - mode| / (L - 1): how concentrated the distribution is at its modal level."""
137
+ L = len(p)
138
+ if L == 1:
139
+ return 1.0
140
+ mode = max(range(L), key=lambda i: p[i])
141
+ return 1.0 - sum(pi * abs(i - mode) for i, pi in enumerate(p)) / (L - 1)
142
+
143
+
144
+ def round_prob(x: float) -> float:
145
+ return round(float(x), 4)
146
+
147
+
148
+ def to_answers(probs, meta) -> dict:
149
+ """Per-question probability lists -> TypeSafe-shaped answers (kev.api.to_answers)."""
150
+ out = {}
151
+ for p, m in zip(probs, meta):
152
+ p = [float(x) for x in p]
153
+ if m["type"] == "noul":
154
+ out[m["id"]] = {"type": "noul", "noul": round_prob(p[1])}
155
+ elif m["type"] == "choice":
156
+ out[m["id"]] = {
157
+ "type": "choice",
158
+ "choice": m["keys"][max(range(len(p)), key=lambda i: p[i])],
159
+ "confidence": round_prob(choice_confidence(p)),
160
+ "probabilities": {k: round_prob(v) for k, v in zip(m["keys"], p)},
161
+ }
162
+ else:
163
+ out[m["id"]] = {
164
+ "type": "score",
165
+ "score": round_prob(sum(i * pi for i, pi in enumerate(p))),
166
+ "confidence": round_prob(score_confidence(p)),
167
+ "legend": m["legend"],
168
+ "probabilities": {str(i): round_prob(v) for i, v in enumerate(p)},
169
+ }
170
+ return out
171
+
172
+
173
+ def output_tokens(tok, answers: dict) -> int:
174
+ """Billing-style figure: tokens of the serialized answers (there is no generation)."""
175
+ return len(tok(json.dumps(answers), add_special_tokens=False).input_ids)
bit_jev/checkpoint.py ADDED
@@ -0,0 +1,149 @@
1
+ """Trained checkpoints: a directory (or Hub repo) holding a LoRA adapter, `head.pt` and the tokenizer.
2
+
3
+ ck = BitJevCheckpoint("runs/bit-jev-2b") # or a hub id
4
+ tok, model = ck.load("cuda", dtype=torch.float32)
5
+
6
+ `head.pt` carries the architecture (Meta) and the pointer-head state dict; the load path is the only
7
+ place that knows both, so train / eval / serve / convert all build the same model.
8
+ """
9
+ import os
10
+ import re
11
+ from dataclasses import dataclass, field
12
+ from pathlib import Path
13
+
14
+ import torch
15
+
16
+ from . import hub as _hub
17
+ from .model import DecisionModel, load_tokenizer
18
+
19
+ HUB_ID = re.compile(r"[\w.-]+/[\w.-]+(@[\w.-]+)?")
20
+
21
+
22
+ def is_hub_id(run):
23
+ return not os.path.isdir(run) and HUB_ID.fullmatch(str(run)) is not None
24
+
25
+
26
+ def resolve_run(run):
27
+ """Local directory as given, or a Hub repo id (optionally @revision) downloaded to the HF cache."""
28
+ if os.path.isdir(run):
29
+ return str(run)
30
+ repo, _, revision = str(run).partition("@")
31
+ return _hub.snapshot(repo, revision=revision or None,
32
+ allow_patterns=["*.json", "*.safetensors", "*.pt", "*.txt", "*.jinja", "*.model"])
33
+
34
+
35
+ @dataclass
36
+ class Meta:
37
+ """Contents of head.pt. `extra` keeps everything else in the file (training args, hashes, fits)."""
38
+ base: str
39
+ head: dict | None = None
40
+ base_revision: str | None = None
41
+ lora: int = 0
42
+ head_dim: int = 256
43
+ temperature: float = 1.0
44
+ arch: str = "bit-jev/0.1"
45
+ holdout: list = field(default_factory=list)
46
+ extra: dict = field(default_factory=dict)
47
+
48
+ KNOWN = ("base", "head", "base_revision", "lora", "head_dim", "temperature", "arch", "holdout")
49
+
50
+ @classmethod
51
+ def from_dict(cls, d):
52
+ if d.get("arch") not in (None, "bit-jev/0.1"):
53
+ raise ValueError(f"head.pt arch {d.get('arch')!r} is not a bit-jev checkpoint")
54
+ return cls(**{k: d[k] for k in cls.KNOWN if k in d},
55
+ extra={k: v for k, v in d.items() if k not in cls.KNOWN})
56
+
57
+ def to_dict(self):
58
+ return {**self.extra, **{k: getattr(self, k) for k in self.KNOWN}}
59
+
60
+
61
+ def read_meta(run):
62
+ return Meta.from_dict(torch.load(f"{run}/head.pt", map_location="cpu", weights_only=False))
63
+
64
+
65
+ def write_meta(run, meta):
66
+ torch.save(meta.to_dict(), f"{run}/head.pt")
67
+
68
+
69
+ class BitJevCheckpoint:
70
+ def __init__(self, run):
71
+ self.requested = str(run)
72
+ self.path = resolve_run(run)
73
+ self.meta = read_meta(self.path)
74
+
75
+ def file(self, name):
76
+ return Path(self.path) / name
77
+
78
+ def load(self, device, dtype=torch.float32, merge=None, temperature=None):
79
+ """-> (tokenizer, model) in eval mode.
80
+
81
+ merge=None defaults by base: on a BitNet base the LoRA adapter must stay UNMERGED. The
82
+ online quantizer (transformers' bitnet integration) rounds every weight to {-1,0,+1} at
83
+ each forward, so a merged delta is quantized away and the served model silently loses
84
+ everything the adapter learned (73.1% unmerged vs 32.4% merged on decision-v7 dev).
85
+ An unmerged LoRA side-branch computes in full precision and reproduces training exactly.
86
+ merge=True remains available for dense fp bases (kev-style), where merging is exact.
87
+
88
+ Distilled runs (bit_jev.distill) store a full backbone instead of an adapter: a run dir
89
+ with config.json and no adapter_model.safetensors is loaded as the backbone itself."""
90
+ has_adapter = self.file("adapter_model.safetensors").exists()
91
+ meta = self.meta
92
+ if has_adapter or not self.file("tokenizer_config.json").exists():
93
+ tok = load_tokenizer(meta.base, revision=meta.base_revision)
94
+ else:
95
+ tok = load_tokenizer(self.path) # distilled run: the run dir carries its own tokenizer
96
+ if has_adapter:
97
+ backbone = meta.base
98
+ else:
99
+ if not self.file("config.json").exists():
100
+ raise FileNotFoundError(f"{self.path} has neither an adapter nor a full backbone "
101
+ "(config.json missing); not a bit-jev run dir")
102
+ backbone = self.path
103
+ m = DecisionModel(backbone, tok, device, lora=None,
104
+ revision=meta.base_revision if has_adapter else None,
105
+ head_dim=meta.head_dim, dtype=dtype)
106
+ if has_adapter:
107
+ from peft import PeftModel
108
+ m.lm = PeftModel.from_pretrained(m.lm, self.path, torch_device=str(device)).to(device)
109
+ if merge is None:
110
+ merge = not self._is_bitnet_base()
111
+ if merge:
112
+ m.lm = m.lm.merge_and_unload()
113
+ if dtype != torch.float32:
114
+ m.lm = m.lm.to(dtype)
115
+ m.head.load_state_dict(meta.head)
116
+ m.head.temperature = meta.temperature if temperature is None else float(temperature)
117
+ m.eval()
118
+ return tok, m
119
+
120
+ def _is_bitnet_base(self):
121
+ """True when the checkpoint's base uses BitNet's online ternary quantization."""
122
+ if "bitnet" in (self.meta.base or "").lower():
123
+ return True
124
+ try:
125
+ from . import hub as _hub
126
+ cfg = _hub.load_config(self.meta.base, revision=self.meta.base_revision)
127
+ return getattr(cfg, "model_type", "") == "bitnet"
128
+ except Exception:
129
+ return False
130
+
131
+ COMPAT_FIELDS = ("base", "lora", "head_dim")
132
+
133
+ def warm_start(self, model, ours: Meta):
134
+ """Delta training: load this checkpoint's adapter and head into a fresh DecisionModel;
135
+ architecture fields are compared before loading (a silent mismatch still "trains" otherwise)."""
136
+ from peft import get_peft_model_state_dict, load_peft_weights, set_peft_model_state_dict
137
+ for name in self.COMPAT_FIELDS:
138
+ if getattr(self.meta, name) != getattr(ours, name):
139
+ raise ValueError(f"--init_from {self.path}: {name} differs "
140
+ f"({getattr(self.meta, name)!r} vs {getattr(ours, name)!r})")
141
+ weights = load_peft_weights(self.path, device="cpu")
142
+ have = set(get_peft_model_state_dict(model.lm))
143
+ if set(weights) != have:
144
+ raise ValueError(f"--init_from {self.path}: adapter tensors mismatch "
145
+ f"(missing {sorted(have - set(weights))[:2]}, "
146
+ f"unexpected {sorted(set(weights) - have)[:2]})")
147
+ set_peft_model_state_dict(model.lm, weights)
148
+ model.head.load_state_dict(self.meta.head)
149
+ return {"init_from": self.requested, "resolved": self.path, "adapter_tensors": len(weights)}