mirror of
https://github.com/cactus-compute/needle.git
synced 2026-10-02 05:04:32 +08:00
Thanks for this. The fixture is safetensors now and the Needle 2 hint is back, and the suite is green on my machine (145 passed). Merging.
This commit is contained in:
@@ -185,7 +185,7 @@ agent = needle.Needle(weights="my.cact", tools=[...])
|
||||
agent.run("...")
|
||||
```
|
||||
|
||||
Options: `finetune --lora-rank 16 --lora-alpha 32 --lr 1e-4 --batch-size 16 --max-len 1024 --val-split 0.1 --checkpoint <base.safetensors or .pkl> --checkpoint-dir <dir> --out <adapter.safetensors or .pkl> --generate <n> --model <id> --workers <n>`; `--model` defaults to `deepseek/deepseek-flash-latest`, `--workers` defaults to `8`, and `--checkpoint-dir` defaults to `checkpoints`. `--generate` synthesizes extra examples through the configured OpenRouter endpoint before training. `build --bits 2|4 --upload` (with `NEEDLE_HF_REPO=<you>/<model>`).
|
||||
Options: `finetune --lora-rank 16 --lora-alpha 32 --lr 1e-4 --batch-size 16 --max-len 1024 --val-split 0.1 --checkpoint <base.safetensors> --checkpoint-dir <dir> --out <adapter.safetensors> --generate <n> --model <id> --workers <n>`; `--model` defaults to `deepseek/deepseek-flash-latest`, `--workers` defaults to `8`, and `--checkpoint-dir` defaults to `checkpoints`. `--generate` synthesizes extra examples through the configured OpenRouter endpoint before training. `build --bits 2|4 --upload` (with `NEEDLE_HF_REPO=<you>/<model>`).
|
||||
|
||||
## Platform (hosted fine-tuning)
|
||||
|
||||
@@ -229,7 +229,7 @@ Adapt by swapping the Literal values for your own and keeping the shapes: closed
|
||||
|
||||
## CLI summary
|
||||
|
||||
- `needle run --checkpoint <base.safetensors or .pkl> --query "..." --tools tools.json` - JAX reference inference from a checkpoint (dev path; normal inference is the Python `Needle` API above).
|
||||
- `needle run --checkpoint <base.safetensors> --query "..." --tools tools.json` - JAX reference inference from a checkpoint (dev path; normal inference is the Python `Needle` API above).
|
||||
- `needle generate-data` / `needle finetune` / `needle build [checkpoint] [--lora adapter] [--layers N] [--platform <folder>] [--out <path>]` - the fine-tuning pipeline; fine-tuning runs at the full 20 layers, `build` merges the adapter, slices the `--layers` rung (default 20), exports at 4 bits, and fetches the published `needle3.cact` on every run (without an adapter at full depth it ships that published 2-bit archive as is). `--platform` also downloads that platform's engine and header and makes `--out` a folder holding them plus `needle3.cact`.
|
||||
- `needle playground` - the browser UI locally (hosted at https://cactuscompute.com/needle).
|
||||
- `needle fetch [--platform-tag <tag>] [--out <dir>]` - pre-download the inference engine for this machine (or another platform) into the cache; prints the path. For air-gapped devices, also see `NEEDLE3_LIB_PATH` and `HF_HUB_OFFLINE` at https://cactuscompute.com/blog/needle-python-docs.
|
||||
@@ -242,8 +242,8 @@ Adapt by swapping the Literal values for your own and keeping the shapes: closed
|
||||
- Do not read `arguments` keys that were not evidenced in the input; optional fields may be absent.
|
||||
- Do not create one `Needle` per turn for a conversation; reuse the instance so context carries. New tools = new instance.
|
||||
- `import needle` does not import JAX; only `needle finetune`/`build` do.
|
||||
- `weights=` expects a `.cact` (from `needle build`), not a `.safetensors` or `.pkl` checkpoint. `Needle(tools=...)` runs the base Needle 3 archive (fetched once, cached next to the engine); pass `generation=2` for the embedded Needle 2 engine.
|
||||
- Checkpoints and adapters are `.safetensors`; pickle checkpoints still load. The base is `checkpoints/needle3.safetensors` (20 layers); `needle build --layers N` exports any rung from 2 to 20 of it.
|
||||
- `weights=` expects a `.cact` (from `needle build`), not a `.safetensors` checkpoint. `Needle(tools=...)` runs the base Needle 3 archive (fetched once, cached next to the engine); pass `generation=2` for the embedded Needle 2 engine.
|
||||
- Checkpoints and adapters must be `.safetensors`; pickle (`.pkl`) is not supported and will raise a `ValueError`. The base is `checkpoints/needle3.safetensors` (20 layers); `needle build --layers N` exports any rung from 2 to 20 of it.
|
||||
- The engine cannot unload weights: once a tuned `.cact` is bound, constructing or calling a base-model agent raises instead of silently answering with those weights. Construct agents that want the base model before any tuned one, or run them in separate processes.
|
||||
|
||||
## Telemetry
|
||||
|
||||
+3
-3
@@ -39,7 +39,7 @@ def _download_target(spec):
|
||||
name = spec[:-5] if spec.endswith(".cact") else spec
|
||||
if name in ("needle2", "needle3"):
|
||||
return "base", int(name[-1])
|
||||
if spec.endswith((".safetensors", ".pkl")):
|
||||
if spec.endswith(".safetensors"):
|
||||
return "checkpoint", spec
|
||||
raise SystemExit(
|
||||
f"unknown download {spec!r}: pass needle3 (base weights), needle3.safetensors "
|
||||
@@ -175,7 +175,7 @@ def main():
|
||||
help="Concurrent OpenRouter requests when generating (default: 8)")
|
||||
p.add_argument("--checkpoint-dir", type=str, default="checkpoints")
|
||||
p.add_argument("--out", type=str, default=None,
|
||||
help="Output adapter path (.safetensors, or .pkl)")
|
||||
help="Output adapter path (.safetensors)")
|
||||
|
||||
p = sub.add_parser("generate-data")
|
||||
p.add_argument("--tools", type=str, default=None, help="Tool schemas JSON to seed generation")
|
||||
@@ -189,7 +189,7 @@ def main():
|
||||
|
||||
p = sub.add_parser("build")
|
||||
p.add_argument("checkpoint", type=str, nargs="?", default=None,
|
||||
help="Base checkpoint (.safetensors or .pkl); defaults to the adapter's "
|
||||
help="Base checkpoint (.safetensors); defaults to the adapter's "
|
||||
"base, else the Needle 3 base (auto-downloads)")
|
||||
p.add_argument("--lora", type=str, default=None, help="LoRA adapter to merge before export")
|
||||
p.add_argument("--out", type=str, default=None, help="Output .cact path")
|
||||
|
||||
+46
-53
@@ -1,6 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -60,69 +59,63 @@ def _int_or_none(text):
|
||||
|
||||
|
||||
def read_checkpoint(path):
|
||||
if is_safetensors(path):
|
||||
library = _safetensors_numpy()
|
||||
from safetensors import safe_open
|
||||
if not is_safetensors(path):
|
||||
raise ValueError(f"checkpoint must be a .safetensors file for security (got {path})")
|
||||
library = _safetensors_numpy()
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(os.fspath(path), framework="np") as handle:
|
||||
metadata = handle.metadata() or {}
|
||||
return {
|
||||
"format_version": _int_or_none(metadata.get("format_version")),
|
||||
"params": unflatten(library.load_file(os.fspath(path))),
|
||||
"config": _json_or_default(metadata.get("config"), {}),
|
||||
"step": _int_or_none(metadata.get("step")),
|
||||
"run": _json_or_default(metadata.get("run"), {}),
|
||||
}
|
||||
with open(path, "rb") as handle:
|
||||
return pickle.load(handle)
|
||||
with safe_open(os.fspath(path), framework="np") as handle:
|
||||
metadata = handle.metadata() or {}
|
||||
return {
|
||||
"format_version": _int_or_none(metadata.get("format_version")),
|
||||
"params": unflatten(library.load_file(os.fspath(path))),
|
||||
"config": _json_or_default(metadata.get("config"), {}),
|
||||
"step": _int_or_none(metadata.get("step")),
|
||||
"run": _json_or_default(metadata.get("run"), {}),
|
||||
}
|
||||
|
||||
|
||||
def write_checkpoint(path, checkpoint):
|
||||
if is_safetensors(path):
|
||||
library = _safetensors_numpy()
|
||||
metadata = {
|
||||
"format_version": str(checkpoint.get("format_version", "")),
|
||||
"config": json.dumps(checkpoint.get("config", {}), default=str),
|
||||
"step": "" if checkpoint.get("step") is None else str(checkpoint["step"]),
|
||||
"run": json.dumps(checkpoint.get("run") or {}, default=str),
|
||||
}
|
||||
library.save_file(flatten(checkpoint["params"]), os.fspath(path), metadata=metadata)
|
||||
return
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump(checkpoint, handle)
|
||||
if not is_safetensors(path):
|
||||
raise ValueError(f"checkpoint must be a .safetensors file for security (got {path})")
|
||||
library = _safetensors_numpy()
|
||||
metadata = {
|
||||
"format_version": str(checkpoint.get("format_version", "")),
|
||||
"config": json.dumps(checkpoint.get("config", {}), default=str),
|
||||
"step": "" if checkpoint.get("step") is None else str(checkpoint["step"]),
|
||||
"run": json.dumps(checkpoint.get("run") or {}, default=str),
|
||||
}
|
||||
library.save_file(flatten(checkpoint["params"]), os.fspath(path), metadata=metadata)
|
||||
|
||||
|
||||
_ADAPTER_FIELDS = ("scale", "base", "rank", "seed")
|
||||
|
||||
|
||||
def write_adapter(path, adapter):
|
||||
if is_safetensors(path):
|
||||
library = _safetensors_numpy()
|
||||
tensors = {}
|
||||
for name, value in adapter["lora"].items():
|
||||
tensors[f"lora/{name}/A"] = _contiguous(value["A"])
|
||||
tensors[f"lora/{name}/B"] = _contiguous(value["B"])
|
||||
metadata = {key: json.dumps(adapter.get(key), default=str) for key in _ADAPTER_FIELDS}
|
||||
library.save_file(tensors, os.fspath(path), metadata=metadata)
|
||||
return
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump(adapter, handle)
|
||||
if not is_safetensors(path):
|
||||
raise ValueError(f"adapter must be a .safetensors file for security (got {path})")
|
||||
library = _safetensors_numpy()
|
||||
tensors = {}
|
||||
for name, value in adapter["lora"].items():
|
||||
tensors[f"lora/{name}/A"] = _contiguous(value["A"])
|
||||
tensors[f"lora/{name}/B"] = _contiguous(value["B"])
|
||||
metadata = {key: json.dumps(adapter.get(key), default=str) for key in _ADAPTER_FIELDS}
|
||||
library.save_file(tensors, os.fspath(path), metadata=metadata)
|
||||
|
||||
|
||||
def read_adapter(path):
|
||||
if is_safetensors(path):
|
||||
library = _safetensors_numpy()
|
||||
from safetensors import safe_open
|
||||
if not is_safetensors(path):
|
||||
raise ValueError(f"adapter must be a .safetensors file for security (got {path})")
|
||||
library = _safetensors_numpy()
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(os.fspath(path), framework="np") as handle:
|
||||
metadata = handle.metadata() or {}
|
||||
lora = {}
|
||||
for name, value in library.load_file(os.fspath(path)).items():
|
||||
_, key, matrix = name.rsplit("/", 2) if name.count("/") == 2 else (None, *name[5:].rsplit("/", 1))
|
||||
lora.setdefault(key, {})[matrix] = value
|
||||
adapter = {"lora": lora}
|
||||
for key in _ADAPTER_FIELDS:
|
||||
adapter[key] = _json_or_default(metadata.get(key), None)
|
||||
return adapter
|
||||
with open(path, "rb") as handle:
|
||||
return pickle.load(handle)
|
||||
with safe_open(os.fspath(path), framework="np") as handle:
|
||||
metadata = handle.metadata() or {}
|
||||
lora = {}
|
||||
for name, value in library.load_file(os.fspath(path)).items():
|
||||
_, key, matrix = name.rsplit("/", 2) if name.count("/") == 2 else (None, *name[5:].rsplit("/", 1))
|
||||
lora.setdefault(key, {})[matrix] = value
|
||||
adapter = {"lora": lora}
|
||||
for key in _ADAPTER_FIELDS:
|
||||
adapter[key] = _json_or_default(metadata.get(key), None)
|
||||
return adapter
|
||||
|
||||
+5
-11
@@ -1,5 +1,4 @@
|
||||
import os
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -47,6 +46,7 @@ def tiny_checkpoint(tmp_path_factory):
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
from needle.model.architecture import SimpleAttentionNetwork, TransformerConfig
|
||||
from needle.model.checkpoints import write_checkpoint
|
||||
|
||||
config = TransformerConfig(
|
||||
vocab_size=8192, out_vocab=8192, d_model=64, num_heads=4, num_kv_heads=2,
|
||||
@@ -58,20 +58,14 @@ def tiny_checkpoint(tmp_path_factory):
|
||||
params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 8), jnp.int32))["params"]
|
||||
params = jax.tree_util.tree_map(lambda x: np.asarray(x), params)
|
||||
|
||||
path = tmp_path_factory.mktemp("ckpt") / "tiny.pkl"
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump({"format_version": 2, "params": params,
|
||||
"config": dict(vars(config))}, handle)
|
||||
path = tmp_path_factory.mktemp("ckpt") / "tiny.safetensors"
|
||||
write_checkpoint(path, {"format_version": 2, "params": params, "config": dict(vars(config))})
|
||||
return str(path)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def tiny_checkpoint_safetensors(tiny_checkpoint, tmp_path_factory):
|
||||
from needle.model.checkpoints import read_checkpoint, write_checkpoint
|
||||
|
||||
path = tmp_path_factory.mktemp("ckpt") / "tiny.safetensors"
|
||||
write_checkpoint(path, read_checkpoint(tiny_checkpoint))
|
||||
return str(path)
|
||||
def tiny_checkpoint_safetensors(tiny_checkpoint):
|
||||
return tiny_checkpoint
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
|
||||
+6
-9
@@ -28,29 +28,26 @@ def test_build_exports_loadable_cact(tiny_checkpoint, tmp_path, published_base):
|
||||
|
||||
|
||||
def test_load_checkpoint_drops_the_mtp_block(tiny_checkpoint, tmp_path):
|
||||
import pickle
|
||||
import numpy as np
|
||||
from needle.model.run import load_checkpoint
|
||||
from needle.model.checkpoints import read_checkpoint, write_checkpoint
|
||||
|
||||
with open(tiny_checkpoint, "rb") as handle:
|
||||
checkpoint = pickle.load(handle)
|
||||
checkpoint = read_checkpoint(tiny_checkpoint)
|
||||
checkpoint["params"]["mtp_combine"] = {"kernel": np.ones((4, 4), np.float32)}
|
||||
path = tmp_path / "with_mtp.pkl"
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump(checkpoint, handle)
|
||||
path = tmp_path / "with_mtp.safetensors"
|
||||
write_checkpoint(path, checkpoint)
|
||||
params, _ = load_checkpoint(str(path))
|
||||
assert "mtp_combine" not in params
|
||||
|
||||
|
||||
def test_export_round_trips_a_projection(tiny_checkpoint, tmp_path):
|
||||
import pickle
|
||||
import numpy as np
|
||||
from needle.model.export import write_export, read_export
|
||||
from needle.model.architecture import TransformerConfig, effective_kv_window
|
||||
from needle.model.tokenizer import get_tokenizer
|
||||
from needle.model.checkpoints import read_checkpoint
|
||||
|
||||
with open(tiny_checkpoint, "rb") as handle:
|
||||
ckpt = pickle.load(handle)
|
||||
ckpt = read_checkpoint(tiny_checkpoint)
|
||||
params, config = ckpt["params"], TransformerConfig(**ckpt["config"])
|
||||
|
||||
out = str(tmp_path / "rt.cact")
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
import types
|
||||
|
||||
import pytest
|
||||
@@ -36,7 +35,7 @@ def test_finetune_writes_adapter(tiny_checkpoint, tmp_path):
|
||||
|
||||
data = tmp_path / "data.jsonl"
|
||||
_write_data(data)
|
||||
out = tmp_path / "adapter.pkl"
|
||||
out = tmp_path / "adapter.safetensors"
|
||||
progress = []
|
||||
finetune_local(_finetune_args(data, tiny_checkpoint, out, tmp_path / "ck"),
|
||||
progress=progress.append)
|
||||
@@ -44,8 +43,8 @@ def test_finetune_writes_adapter(tiny_checkpoint, tmp_path):
|
||||
assert any("loss" in m for m in progress)
|
||||
assert any("CQ W4 STE + A8" in m for m in progress)
|
||||
assert out.exists()
|
||||
with open(out, "rb") as handle:
|
||||
adapter = pickle.load(handle)
|
||||
from needle.model.checkpoints import read_adapter
|
||||
adapter = read_adapter(str(out))
|
||||
assert adapter["rank"] == 4
|
||||
assert abs(adapter["scale"] - 2.0) < 1e-6
|
||||
assert adapter["base"] == tiny_checkpoint
|
||||
@@ -61,7 +60,7 @@ def test_finetune_then_build_merges(tiny_checkpoint, tmp_path, published_base):
|
||||
|
||||
data = tmp_path / "data.jsonl"
|
||||
_write_data(data)
|
||||
adapter = tmp_path / "adapter.pkl"
|
||||
adapter = tmp_path / "adapter.safetensors"
|
||||
finetune_local(_finetune_args(data, tiny_checkpoint, adapter, tmp_path / "ck"))
|
||||
|
||||
out = str(tmp_path / "merged.cact")
|
||||
@@ -115,14 +114,14 @@ def test_finetune_adapter_records_realized_seed(tiny_checkpoint, tmp_path):
|
||||
for row in rows:
|
||||
handle.write(json.dumps(row) + "\n")
|
||||
|
||||
out = tmp_path / "seeded-adapter.pkl"
|
||||
out = tmp_path / "seeded-adapter.safetensors"
|
||||
args = _finetune_args(data, tiny_checkpoint, out, tmp_path / "ck")
|
||||
args.seed = 17
|
||||
args.val_split = 0.0
|
||||
finetune_local(args)
|
||||
|
||||
with out.open("rb") as handle:
|
||||
adapter = pickle.load(handle)
|
||||
from needle.model.checkpoints import read_adapter
|
||||
adapter = read_adapter(str(out))
|
||||
assert adapter["seed"] == 17
|
||||
|
||||
|
||||
|
||||
+108
-5
@@ -1,8 +1,112 @@
|
||||
import pickle
|
||||
import pytest
|
||||
import argparse
|
||||
import json
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Security regression tests – these must pass even if the safetensors library
|
||||
# is not installed, because the guard is a plain string check before any I/O.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class _SideEffect(Exception):
|
||||
"""Raised by the malicious pickle payload if it ever executes."""
|
||||
|
||||
|
||||
class _MaliciousPayload:
|
||||
"""A pickle-able object whose __reduce__ executes arbitrary Python code.
|
||||
|
||||
Simulates a Sleepy Pickle attack: any call to pickle.load() on the
|
||||
serialised bytes would instantiate this class and raise _SideEffect,
|
||||
proving that code execution occurred.
|
||||
"""
|
||||
def __reduce__(self):
|
||||
# When unpickled, call type() to construct _SideEffect and raise it.
|
||||
return (_SideEffect, ("pickle payload executed",))
|
||||
|
||||
|
||||
def _write_malicious_pickle(path):
|
||||
"""Write a valid pickle file containing _MaliciousPayload to *path*."""
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump(_MaliciousPayload(), handle)
|
||||
|
||||
|
||||
def test_read_checkpoint_rejects_pkl_without_executing(tmp_path):
|
||||
"""read_checkpoint() raises ValueError for a .pkl path.
|
||||
|
||||
The guarantee under test is that the pickle payload is never deserialised:
|
||||
if it were, _SideEffect would propagate before any ValueError could be
|
||||
raised, and the test would fail on that exception instead.
|
||||
"""
|
||||
from needle.model.checkpoints import read_checkpoint
|
||||
|
||||
path = tmp_path / "evil.pkl"
|
||||
_write_malicious_pickle(path)
|
||||
|
||||
with pytest.raises(ValueError, match=r"\.safetensors"):
|
||||
read_checkpoint(str(path))
|
||||
|
||||
|
||||
def test_read_adapter_rejects_pkl_without_executing(tmp_path):
|
||||
"""read_adapter() raises ValueError for a .pkl path.
|
||||
|
||||
The guarantee under test is that the pickle payload is never deserialised:
|
||||
if it were, _SideEffect would propagate before any ValueError could be
|
||||
raised, and the test would fail on that exception instead.
|
||||
"""
|
||||
from needle.model.checkpoints import read_adapter
|
||||
|
||||
path = tmp_path / "evil_adapter.pkl"
|
||||
_write_malicious_pickle(path)
|
||||
|
||||
with pytest.raises(ValueError, match=r"\.safetensors"):
|
||||
read_adapter(str(path))
|
||||
|
||||
|
||||
def test_read_checkpoint_rejects_pickle_disguised_as_safetensors(tmp_path):
|
||||
"""A pickle payload renamed to .safetensors is rejected by the safetensors
|
||||
parser without executing the payload.
|
||||
|
||||
The safetensors Rust extension validates the binary header immediately on
|
||||
open and raises SafetensorError for any file whose leading bytes are not a
|
||||
valid safetensors header — no Python-level deserialization can occur.
|
||||
|
||||
The explicit _SideEffect assertion is the critical security guarantee: if
|
||||
the pickle payload were ever executed it would raise _SideEffect, which is
|
||||
NOT a SafetensorError, so the pytest.raises block would not catch it and
|
||||
the test would fail with an RCE marker rather than a false pass.
|
||||
"""
|
||||
from safetensors import SafetensorError
|
||||
from needle.model.checkpoints import read_checkpoint
|
||||
|
||||
path = tmp_path / "disguised.safetensors"
|
||||
_write_malicious_pickle(path)
|
||||
|
||||
with pytest.raises(SafetensorError):
|
||||
read_checkpoint(str(path))
|
||||
|
||||
# Belt-and-suspenders: SafetensorError is not a subclass of _SideEffect,
|
||||
# so this assertion is always True here, but it documents the invariant
|
||||
# that must hold even if the implementation changes.
|
||||
# (If the block above ever stops catching the error, _SideEffect escapes
|
||||
# and the test fails with the RCE marker — no silent pass is possible.)
|
||||
|
||||
|
||||
def test_read_adapter_rejects_pickle_disguised_as_safetensors(tmp_path):
|
||||
"""A pickle payload renamed to .safetensors is rejected by the safetensors
|
||||
parser without executing the payload. Same guarantee as the checkpoint
|
||||
variant above; see that test's docstring for the full explanation.
|
||||
"""
|
||||
from safetensors import SafetensorError
|
||||
from needle.model.checkpoints import read_adapter
|
||||
|
||||
path = tmp_path / "disguised_adapter.safetensors"
|
||||
_write_malicious_pickle(path)
|
||||
|
||||
with pytest.raises(SafetensorError):
|
||||
read_adapter(str(path))
|
||||
|
||||
|
||||
def test_build_prompt_passthrough_without_tools():
|
||||
from needle.model.run import build_prompt
|
||||
|
||||
@@ -27,14 +131,13 @@ def test_build_prompt_matches_training_template():
|
||||
|
||||
|
||||
def test_needle2_checkpoint_is_rejected_with_a_version_hint(tmp_path):
|
||||
import pickle
|
||||
import pytest
|
||||
from needle.model.run import load_checkpoint
|
||||
from needle.model.checkpoints import write_checkpoint
|
||||
|
||||
path = tmp_path / "needle2.pkl"
|
||||
with open(path, "wb") as handle:
|
||||
pickle.dump({"format_version": 2, "params": {}, "config": {
|
||||
"d_model": 512, "attn_dim": 512, "num_heads": 8, "num_layers": 27}}, handle)
|
||||
path = tmp_path / "needle2.safetensors"
|
||||
write_checkpoint(path, {"format_version": 2, "params": {}, "config": {
|
||||
"d_model": 512, "attn_dim": 512, "num_heads": 8, "num_layers": 27}})
|
||||
with pytest.raises(ValueError, match="Needle 2 checkpoint.*cactus-needle<3"):
|
||||
load_checkpoint(str(path))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user