revisionlab 0.1.0__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.
- revisionlab/__init__.py +19 -0
- revisionlab/baselines.py +148 -0
- revisionlab/diagnostics.py +42 -0
- revisionlab/memory.py +195 -0
- revisionlab/protocols.py +123 -0
- revisionlab/py.typed +0 -0
- revisionlab/replay.py +204 -0
- revisionlab/serialization.py +113 -0
- revisionlab/smoke.py +97 -0
- revisionlab-0.1.0.dist-info/METADATA +109 -0
- revisionlab-0.1.0.dist-info/RECORD +14 -0
- revisionlab-0.1.0.dist-info/WHEEL +4 -0
- revisionlab-0.1.0.dist-info/entry_points.txt +2 -0
- revisionlab-0.1.0.dist-info/licenses/LICENSE +21 -0
revisionlab/__init__.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
"""Delayed evidence and residual-memory experiments, native to PyTorch."""
|
|
2
|
+
|
|
3
|
+
from .diagnostics import aggregate_alignment, credit_alignment
|
|
4
|
+
from .memory import ResidualMemory, SparseResidualMemory, delta_write
|
|
5
|
+
from .protocols import BudgetReport, ForecastTicket, MemoryProtocol, ReleasedEvidence
|
|
6
|
+
|
|
7
|
+
__version__ = "0.1.0"
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"BudgetReport",
|
|
11
|
+
"ForecastTicket",
|
|
12
|
+
"MemoryProtocol",
|
|
13
|
+
"ReleasedEvidence",
|
|
14
|
+
"ResidualMemory",
|
|
15
|
+
"SparseResidualMemory",
|
|
16
|
+
"aggregate_alignment",
|
|
17
|
+
"credit_alignment",
|
|
18
|
+
"delta_write",
|
|
19
|
+
]
|
revisionlab/baselines.py
ADDED
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
"""Existing recursive least squares and explicit negative-control baselines."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from dataclasses import replace
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import torch
|
|
8
|
+
from numpy.typing import NDArray
|
|
9
|
+
from torch import Tensor, nn
|
|
10
|
+
|
|
11
|
+
from .memory import ResidualMemory
|
|
12
|
+
from .protocols import ReleasedEvidence
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _key(values: Tensor, key: Tensor | NDArray[np.float64]) -> Tensor:
|
|
16
|
+
if isinstance(key, Tensor):
|
|
17
|
+
if key.dtype != values.dtype or key.device != values.device:
|
|
18
|
+
raise ValueError("Tensor keys must match memory dtype and device")
|
|
19
|
+
result = key
|
|
20
|
+
else:
|
|
21
|
+
result = torch.as_tensor(np.array(key, copy=True), dtype=values.dtype, device=values.device)
|
|
22
|
+
if result.shape != values.shape or not bool(torch.isfinite(result).all()):
|
|
23
|
+
raise ValueError("Key must be a finite vector matching memory width")
|
|
24
|
+
return result
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
class RLS_Corrector(nn.Module): # type: ignore[misc]
|
|
28
|
+
"""FP64-default RLS residual corrector, reproducing the v13 baseline.
|
|
29
|
+
|
|
30
|
+
Initial covariance is 100 I. Forgetting lambda is in (0,1]; observed
|
|
31
|
+
residuals are clipped to [-12,12]. Low-level writes do not deduplicate:
|
|
32
|
+
use DelayedReplay for owner-bound, exactly-once evidence handling.
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
def __init__(self, slots: int, forget: float = 0.999, *, dtype: torch.dtype = torch.float64):
|
|
36
|
+
super().__init__()
|
|
37
|
+
if isinstance(slots, bool) or not isinstance(slots, int) or slots < 1:
|
|
38
|
+
raise ValueError("slots must be a positive integer")
|
|
39
|
+
if not math.isfinite(forget) or not 0 < forget <= 1:
|
|
40
|
+
raise ValueError("forget must be finite and in (0,1]")
|
|
41
|
+
if dtype not in (torch.float32, torch.float64):
|
|
42
|
+
raise ValueError("dtype must be float32 or float64")
|
|
43
|
+
self.register_buffer("forgetting_factor", torch.tensor(forget, dtype=dtype))
|
|
44
|
+
self.register_buffer("values", torch.zeros(slots, dtype=dtype))
|
|
45
|
+
self.register_buffer("P", 100 * torch.eye(slots, dtype=dtype))
|
|
46
|
+
self.register_buffer("writes", torch.zeros((), dtype=torch.int64))
|
|
47
|
+
self.register_buffer("last_label_hour", torch.full((), -1, dtype=torch.int64))
|
|
48
|
+
|
|
49
|
+
def read(self, key: Tensor | NDArray[np.float64]) -> Tensor:
|
|
50
|
+
result = torch.dot(self.values, _key(self.values, key))
|
|
51
|
+
if not bool(torch.isfinite(result)):
|
|
52
|
+
raise FloatingPointError("RLS read is non-finite")
|
|
53
|
+
return result
|
|
54
|
+
|
|
55
|
+
@property
|
|
56
|
+
def forget(self) -> float:
|
|
57
|
+
return float(self.forgetting_factor)
|
|
58
|
+
|
|
59
|
+
@torch.no_grad() # type: ignore[untyped-decorator]
|
|
60
|
+
def write(self, evidence: ReleasedEvidence) -> None:
|
|
61
|
+
if evidence.ticket.address_version != 0:
|
|
62
|
+
raise ValueError("Only address version 0 is supported")
|
|
63
|
+
if int(self.writes) >= torch.iinfo(torch.int64).max:
|
|
64
|
+
raise OverflowError("RLS write counter exhausted; write not committed")
|
|
65
|
+
key = _key(self.values, evidence.ticket.read_key)
|
|
66
|
+
residual = evidence.target - evidence.ticket.base_prediction
|
|
67
|
+
if not math.isfinite(residual):
|
|
68
|
+
raise ValueError("Residual must be finite")
|
|
69
|
+
target = max(-12.0, min(12.0, residual))
|
|
70
|
+
px = self.P @ key
|
|
71
|
+
denominator = self.forget + torch.dot(key, px)
|
|
72
|
+
if not bool(torch.isfinite(denominator)) or float(denominator) <= 0:
|
|
73
|
+
raise FloatingPointError("RLS denominator is not finite and positive")
|
|
74
|
+
gain = px / denominator
|
|
75
|
+
values = self.values + gain * (target - self.read(key))
|
|
76
|
+
covariance = (self.P - torch.outer(gain, px)) / self.forget
|
|
77
|
+
covariance = (covariance + covariance.T) * 0.5
|
|
78
|
+
if not bool(torch.isfinite(values).all() and torch.isfinite(covariance).all()):
|
|
79
|
+
raise FloatingPointError("RLS candidate state is non-finite; write not committed")
|
|
80
|
+
self.values.copy_(values)
|
|
81
|
+
self.P.copy_(covariance)
|
|
82
|
+
self.writes.add_(1)
|
|
83
|
+
self.last_label_hour.fill_(evidence.ticket.available_hour)
|
|
84
|
+
|
|
85
|
+
def config(self) -> dict[str, object]:
|
|
86
|
+
return {
|
|
87
|
+
"slots": self.values.numel(),
|
|
88
|
+
"forget": self.forget,
|
|
89
|
+
"dtype": str(self.values.dtype).removeprefix("torch."),
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
def snapshot(self) -> dict[str, object]:
|
|
93
|
+
return {
|
|
94
|
+
"values": self.values.detach().clone(),
|
|
95
|
+
"P": self.P.detach().clone(),
|
|
96
|
+
"writes": int(self.writes),
|
|
97
|
+
"last_label_hour": None if int(self.writes) == 0 else int(self.last_label_hour),
|
|
98
|
+
"forget": self.forget,
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
class NoWrite(ResidualMemory):
|
|
103
|
+
"""Zero-correction control: accept mature evidence but never modify state."""
|
|
104
|
+
|
|
105
|
+
def __init__(self, slots: int, *, dtype: torch.dtype = torch.float64):
|
|
106
|
+
super().__init__(slots, 1.0, dtype=dtype)
|
|
107
|
+
|
|
108
|
+
def read(self, key: Tensor | NDArray[np.float64]) -> Tensor:
|
|
109
|
+
_key(self.values, key)
|
|
110
|
+
return torch.zeros((), dtype=self.values.dtype, device=self.values.device)
|
|
111
|
+
|
|
112
|
+
def write(self, evidence: ReleasedEvidence) -> None:
|
|
113
|
+
if evidence.ticket.address_version != 0:
|
|
114
|
+
raise ValueError("Only address version 0 is supported")
|
|
115
|
+
_key(self.values, evidence.ticket.read_key)
|
|
116
|
+
if not math.isfinite(evidence.target - evidence.ticket.base_prediction):
|
|
117
|
+
raise ValueError("Residual must be finite")
|
|
118
|
+
|
|
119
|
+
def config(self) -> dict[str, object]:
|
|
120
|
+
return {
|
|
121
|
+
"slots": self.values.numel(),
|
|
122
|
+
"dtype": str(self.values.dtype).removeprefix("torch."),
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
class ShuffledWrite(ResidualMemory):
|
|
127
|
+
"""Negative control: permute write addresses, keep read addresses unchanged."""
|
|
128
|
+
|
|
129
|
+
def __init__(
|
|
130
|
+
self, slots: int, rate: float = 0.05, *, seed: int = 0, dtype: torch.dtype = torch.float64
|
|
131
|
+
):
|
|
132
|
+
if slots < 2:
|
|
133
|
+
raise ValueError("Shuffled-write requires at least two slots")
|
|
134
|
+
super().__init__(slots, rate, dtype=dtype)
|
|
135
|
+
self.seed = seed
|
|
136
|
+
generator = torch.Generator().manual_seed(seed)
|
|
137
|
+
permutation = torch.randperm(slots, generator=generator)
|
|
138
|
+
if torch.equal(permutation, torch.arange(slots)):
|
|
139
|
+
permutation = permutation.roll(1)
|
|
140
|
+
self.register_buffer("permutation", permutation)
|
|
141
|
+
|
|
142
|
+
def write(self, evidence: ReleasedEvidence) -> None:
|
|
143
|
+
key = _key(self.values, evidence.ticket.read_key)
|
|
144
|
+
shuffled = key[self.permutation].detach().cpu().numpy()
|
|
145
|
+
super().write(replace(evidence, ticket=replace(evidence.ticket, read_key=shuffled)))
|
|
146
|
+
|
|
147
|
+
def config(self) -> dict[str, object]:
|
|
148
|
+
return {**super().config(), "seed": self.seed}
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""Single-trial cosine and pooled rho; zero-norm alignment is undefined."""
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
from numpy.typing import NDArray
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _array(value: Any) -> NDArray[np.float64]:
|
|
11
|
+
if isinstance(value, torch.Tensor):
|
|
12
|
+
value = value.detach().to(device="cpu", dtype=torch.float64).numpy()
|
|
13
|
+
return np.asarray(value, dtype=np.float64)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _alignment(first: Any, second: Any, *, flatten: bool) -> float | None:
|
|
17
|
+
a, b = _array(first), _array(second)
|
|
18
|
+
if flatten:
|
|
19
|
+
a, b = a.ravel(), b.ravel()
|
|
20
|
+
if a.shape != b.shape or not (np.isfinite(a).all() and np.isfinite(b).all()):
|
|
21
|
+
raise ValueError("Finite matched vectors required")
|
|
22
|
+
if a.size == 0:
|
|
23
|
+
return None
|
|
24
|
+
scale_a, scale_b = float(np.max(np.abs(a))), float(np.max(np.abs(b)))
|
|
25
|
+
if scale_a == 0 or scale_b == 0:
|
|
26
|
+
return None
|
|
27
|
+
# Independent global scaling leaves pooled rho unchanged, avoids over/underflow.
|
|
28
|
+
a, b = a / scale_a, b / scale_b
|
|
29
|
+
denominator = np.sqrt(np.sum(a * a)) * np.sqrt(np.sum(b * b))
|
|
30
|
+
return float(np.clip(np.sum(a * b) / denominator, -1, 1))
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def credit_alignment(update: Any, negative_gradient: Any) -> dict[str, bool | float | None]:
|
|
34
|
+
"""Cosine of an actual update with a reference negative gradient."""
|
|
35
|
+
cosine = _alignment(update, negative_gradient, flatten=True)
|
|
36
|
+
return {"defined": cosine is not None, "cosine": cosine}
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def aggregate_alignment(updates: Any, negative_gradients: Any) -> dict[str, bool | float | None]:
|
|
40
|
+
"""sum dot / sqrt(sum squared norms), not an average of trial cosines."""
|
|
41
|
+
rho = _alignment(updates, negative_gradients, flatten=False)
|
|
42
|
+
return {"defined": rho is not None, "rho": rho}
|
revisionlab/memory.py
ADDED
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
"""Normalized delta/LMS updates; evidence ownership belongs to the replay runner."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from numbers import Integral
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import torch
|
|
9
|
+
from numpy.typing import NDArray
|
|
10
|
+
from torch import Tensor, nn
|
|
11
|
+
|
|
12
|
+
from .protocols import ReleasedEvidence, _key
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _size_rate(slots: int, rate: float) -> None:
|
|
16
|
+
if isinstance(slots, bool) or not isinstance(slots, Integral) or slots < 1:
|
|
17
|
+
raise ValueError("slots must be a positive integer")
|
|
18
|
+
if isinstance(rate, bool) or not math.isfinite(rate) or not 0 < rate <= 1:
|
|
19
|
+
raise ValueError("rate must be finite and in (0, 1]")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def delta_write(values: Tensor, key: Tensor, target: Tensor, rate: float) -> Tensor:
|
|
23
|
+
"""Pure differentiable update: [...,slots,channels], [...,slots], [...,channels].
|
|
24
|
+
|
|
25
|
+
All tensors must share dtype/device and exactly matching leading dimensions.
|
|
26
|
+
Finite inputs and representable intermediate squared key norms are required;
|
|
27
|
+
stateful backends validate norms and proposals before committing. This
|
|
28
|
+
function neither clips residuals nor mutates inputs.
|
|
29
|
+
"""
|
|
30
|
+
if values.ndim < 2 or key.ndim != values.ndim - 1 or target.ndim != values.ndim - 1:
|
|
31
|
+
raise ValueError(
|
|
32
|
+
"Expected values [...,slots,channels], key [...,slots], target [...,channels]"
|
|
33
|
+
)
|
|
34
|
+
if (
|
|
35
|
+
values.shape[-2] < 1
|
|
36
|
+
or values.shape[-1] < 1
|
|
37
|
+
or key.shape != values.shape[:-1]
|
|
38
|
+
or target.shape != (*values.shape[:-2], values.shape[-1])
|
|
39
|
+
):
|
|
40
|
+
raise ValueError("Delta update shape mismatch")
|
|
41
|
+
if not values.is_floating_point() or values.dtype != key.dtype or values.dtype != target.dtype:
|
|
42
|
+
raise ValueError("Delta update requires one shared floating dtype")
|
|
43
|
+
if values.device != key.device or values.device != target.device:
|
|
44
|
+
raise ValueError("Delta update requires one shared device")
|
|
45
|
+
if isinstance(rate, bool) or not math.isfinite(rate) or not 0 < rate <= 1:
|
|
46
|
+
raise ValueError("rate must be finite and in (0, 1]")
|
|
47
|
+
prediction = torch.einsum("...n,...nd->...d", key, values)
|
|
48
|
+
floor = max(1e-12, torch.finfo(values.dtype).tiny)
|
|
49
|
+
denominator = key.square().sum(-1, keepdim=True).clamp_min(floor)
|
|
50
|
+
return values + (rate / denominator).unsqueeze(-1) * key.unsqueeze(-1) * (
|
|
51
|
+
target - prediction
|
|
52
|
+
).unsqueeze(-2)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class SparseResidualMemory:
|
|
56
|
+
"""Backward-compatible NumPy scalar backend; one instance per horizon."""
|
|
57
|
+
|
|
58
|
+
def __init__(
|
|
59
|
+
self, slots: int, rate: float, permutation: NDArray[np.int64] | None = None
|
|
60
|
+
) -> None:
|
|
61
|
+
_size_rate(slots, rate)
|
|
62
|
+
self.values: NDArray[np.float64] = np.zeros(slots, dtype=np.float64)
|
|
63
|
+
self.rate = float(rate)
|
|
64
|
+
self.writes = 0
|
|
65
|
+
self.last_label_hour: int | None = None
|
|
66
|
+
self.permutation: NDArray[np.int64] | None = None
|
|
67
|
+
if permutation is not None:
|
|
68
|
+
candidate = np.asarray(permutation)
|
|
69
|
+
if (
|
|
70
|
+
candidate.shape != (slots,)
|
|
71
|
+
or not np.issubdtype(candidate.dtype, np.integer)
|
|
72
|
+
or not np.array_equal(np.sort(candidate), np.arange(slots))
|
|
73
|
+
):
|
|
74
|
+
raise ValueError("Permutation must be an integral bijection")
|
|
75
|
+
self.permutation = np.frombuffer(candidate.astype(np.int64).tobytes(), dtype=np.int64)
|
|
76
|
+
|
|
77
|
+
def read(self, key: Any) -> float:
|
|
78
|
+
result = float(np.dot(self.values, _key(key, self.values.size)))
|
|
79
|
+
if not math.isfinite(result):
|
|
80
|
+
raise FloatingPointError("Memory read is not finite")
|
|
81
|
+
return result
|
|
82
|
+
|
|
83
|
+
def write(self, evidence: ReleasedEvidence) -> None:
|
|
84
|
+
if evidence.ticket.address_version != 0:
|
|
85
|
+
raise ValueError("Unsupported address_version")
|
|
86
|
+
key = _key(evidence.ticket.read_key, self.values.size)
|
|
87
|
+
if self.permutation is not None:
|
|
88
|
+
key = key[self.permutation]
|
|
89
|
+
with np.errstate(over="ignore", invalid="ignore"):
|
|
90
|
+
squared_norm = float(key @ key)
|
|
91
|
+
if not math.isfinite(squared_norm):
|
|
92
|
+
raise FloatingPointError("Key squared norm is not finite")
|
|
93
|
+
residual = float(np.clip(evidence.target - evidence.ticket.base_prediction, -12, 12))
|
|
94
|
+
with np.errstate(over="ignore", invalid="ignore", divide="ignore"):
|
|
95
|
+
proposal = self.values + self.rate * key * (residual - np.dot(self.values, key)) / max(
|
|
96
|
+
squared_norm, 1e-12
|
|
97
|
+
)
|
|
98
|
+
if not np.isfinite(proposal).all():
|
|
99
|
+
raise FloatingPointError("Memory proposal is not finite")
|
|
100
|
+
self.values[:] = proposal
|
|
101
|
+
self.writes += 1
|
|
102
|
+
self.last_label_hour = evidence.ticket.available_hour
|
|
103
|
+
|
|
104
|
+
def snapshot(self) -> dict[str, Any]:
|
|
105
|
+
return {
|
|
106
|
+
"values": self.values.copy(),
|
|
107
|
+
"writes": self.writes,
|
|
108
|
+
"rate": self.rate,
|
|
109
|
+
"last_label_hour": self.last_label_hour,
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
class ResidualMemory(nn.Module): # type: ignore[misc]
|
|
114
|
+
"""Native Torch scalar state; one instance per horizon, default FP64.
|
|
115
|
+
|
|
116
|
+
Tensor keys require an exact dtype/device match. NumPy ticket keys are
|
|
117
|
+
explicitly copied into this backend's configured dtype/device. Stateful
|
|
118
|
+
writes use clipped target-minus-base residuals and have no autograd graph.
|
|
119
|
+
"""
|
|
120
|
+
|
|
121
|
+
def __init__(
|
|
122
|
+
self,
|
|
123
|
+
slots: int,
|
|
124
|
+
rate: float,
|
|
125
|
+
*,
|
|
126
|
+
dtype: torch.dtype = torch.float64,
|
|
127
|
+
device: torch.device | str | None = None,
|
|
128
|
+
) -> None:
|
|
129
|
+
super().__init__()
|
|
130
|
+
_size_rate(slots, rate)
|
|
131
|
+
if not dtype.is_floating_point:
|
|
132
|
+
raise ValueError("Memory requires a floating dtype")
|
|
133
|
+
self.register_buffer("values", torch.zeros(slots, dtype=dtype, device=device))
|
|
134
|
+
self.register_buffer("learning_rate", torch.tensor(rate, dtype=dtype, device=device))
|
|
135
|
+
self.register_buffer("writes", torch.zeros((), dtype=torch.int64, device=device))
|
|
136
|
+
self.register_buffer(
|
|
137
|
+
"last_label_hour", torch.full((), -1, dtype=torch.int64, device=device)
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
@property
|
|
141
|
+
def rate(self) -> float:
|
|
142
|
+
return float(self.learning_rate.item())
|
|
143
|
+
|
|
144
|
+
def _read_key(self, key: Any) -> Tensor:
|
|
145
|
+
if isinstance(key, Tensor):
|
|
146
|
+
if key.dtype != self.values.dtype or key.device != self.values.device:
|
|
147
|
+
raise ValueError("Tensor key dtype/device must match memory")
|
|
148
|
+
result = key
|
|
149
|
+
else:
|
|
150
|
+
result = torch.tensor(_key(key), dtype=self.values.dtype, device=self.values.device)
|
|
151
|
+
if result.shape != self.values.shape or not torch.isfinite(result).all():
|
|
152
|
+
raise ValueError("A finite key with matching address width is required")
|
|
153
|
+
return result
|
|
154
|
+
|
|
155
|
+
def read(self, key: Any) -> Tensor:
|
|
156
|
+
result = torch.dot(self.values, self._read_key(key))
|
|
157
|
+
if not torch.isfinite(result):
|
|
158
|
+
raise FloatingPointError("Memory read is not finite")
|
|
159
|
+
return result
|
|
160
|
+
|
|
161
|
+
def write(self, evidence: ReleasedEvidence) -> None:
|
|
162
|
+
if evidence.ticket.address_version != 0:
|
|
163
|
+
raise ValueError("Unsupported address_version")
|
|
164
|
+
hour = evidence.ticket.available_hour
|
|
165
|
+
limits = torch.iinfo(torch.int64)
|
|
166
|
+
if not limits.min <= hour <= limits.max or int(self.writes.item()) >= limits.max:
|
|
167
|
+
raise ValueError("Evidence hour/write count exceeds native metadata range")
|
|
168
|
+
key = self._read_key(evidence.ticket.read_key)
|
|
169
|
+
if not torch.isfinite(key.square().sum()):
|
|
170
|
+
raise FloatingPointError("Key squared norm is not finite")
|
|
171
|
+
residual = max(-12.0, min(12.0, evidence.target - evidence.ticket.base_prediction))
|
|
172
|
+
target = self.values.new_tensor([residual])
|
|
173
|
+
with torch.no_grad():
|
|
174
|
+
proposal = delta_write(self.values[:, None], key, target, self.rate).squeeze(-1)
|
|
175
|
+
if not torch.isfinite(proposal).all():
|
|
176
|
+
raise FloatingPointError("Memory proposal is not finite")
|
|
177
|
+
self.values.copy_(proposal)
|
|
178
|
+
self.writes.add_(1)
|
|
179
|
+
self.last_label_hour.fill_(hour)
|
|
180
|
+
|
|
181
|
+
def snapshot(self) -> dict[str, Any]:
|
|
182
|
+
hour = int(self.last_label_hour.item())
|
|
183
|
+
return {
|
|
184
|
+
"values": self.values.detach().clone(),
|
|
185
|
+
"writes": int(self.writes.item()),
|
|
186
|
+
"rate": self.rate,
|
|
187
|
+
"last_label_hour": None if int(self.writes.item()) == 0 else hour,
|
|
188
|
+
}
|
|
189
|
+
|
|
190
|
+
def config(self) -> dict[str, Any]:
|
|
191
|
+
return {
|
|
192
|
+
"slots": self.values.numel(),
|
|
193
|
+
"rate": self.rate,
|
|
194
|
+
"dtype": str(self.values.dtype).removeprefix("torch."),
|
|
195
|
+
}
|
revisionlab/protocols.py
ADDED
|
@@ -0,0 +1,123 @@
|
|
|
1
|
+
"""Immutable issued forecasts and causally released evidence."""
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
from numbers import Integral
|
|
5
|
+
from typing import Any, Protocol, runtime_checkable
|
|
6
|
+
from uuid import uuid4
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
import torch
|
|
10
|
+
from numpy.typing import NDArray
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _tick(value: int, name: str) -> int:
|
|
14
|
+
if isinstance(value, bool) or not isinstance(value, Integral):
|
|
15
|
+
raise ValueError(f"{name} must be an integral tick")
|
|
16
|
+
if not -(2**63) <= value < 2**63:
|
|
17
|
+
raise ValueError(f"{name} must fit the signed int64 tick range")
|
|
18
|
+
return int(value)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _scalar(value: float, name: str) -> float:
|
|
22
|
+
if isinstance(value, bool):
|
|
23
|
+
raise ValueError(f"{name} must be a finite scalar")
|
|
24
|
+
if isinstance(value, torch.Tensor):
|
|
25
|
+
value = value.detach().cpu().numpy()
|
|
26
|
+
array = np.asarray(value)
|
|
27
|
+
if array.ndim != 0:
|
|
28
|
+
raise ValueError(f"{name} must be a finite scalar")
|
|
29
|
+
try:
|
|
30
|
+
result = float(array)
|
|
31
|
+
except (TypeError, ValueError) as error:
|
|
32
|
+
raise ValueError(f"{name} must be a finite scalar") from error
|
|
33
|
+
if not np.isfinite(result):
|
|
34
|
+
raise ValueError(f"{name} must be a finite scalar")
|
|
35
|
+
return result
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _key(value: Any, slots: int | None = None) -> NDArray[np.float64]:
|
|
39
|
+
if isinstance(value, torch.Tensor):
|
|
40
|
+
value = value.detach().to(device="cpu", dtype=torch.float64).numpy()
|
|
41
|
+
result = np.asarray(value, dtype=np.float64)
|
|
42
|
+
if result.ndim != 1 or result.size == 0 or not np.isfinite(result).all():
|
|
43
|
+
raise ValueError("A finite nonempty one-dimensional key is required")
|
|
44
|
+
if slots is not None and result.size != slots:
|
|
45
|
+
raise ValueError("Address width mismatch")
|
|
46
|
+
return result
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclass(frozen=True)
|
|
50
|
+
class ForecastTicket:
|
|
51
|
+
issue_hour: int
|
|
52
|
+
horizon: int
|
|
53
|
+
base_prediction: float
|
|
54
|
+
read_key: NDArray[np.float64]
|
|
55
|
+
prediction: float | None = None
|
|
56
|
+
ticket_id: str = field(default_factory=lambda: str(uuid4()))
|
|
57
|
+
address_version: int = 0
|
|
58
|
+
|
|
59
|
+
def __post_init__(self) -> None:
|
|
60
|
+
object.__setattr__(self, "issue_hour", _tick(self.issue_hour, "issue_hour"))
|
|
61
|
+
horizon = _tick(self.horizon, "horizon")
|
|
62
|
+
if horizon < 1:
|
|
63
|
+
raise ValueError("Horizon must be positive")
|
|
64
|
+
object.__setattr__(self, "horizon", horizon)
|
|
65
|
+
_tick(self.issue_hour + horizon, "available_hour")
|
|
66
|
+
object.__setattr__(
|
|
67
|
+
self, "base_prediction", _scalar(self.base_prediction, "base_prediction")
|
|
68
|
+
)
|
|
69
|
+
if self.prediction is not None:
|
|
70
|
+
object.__setattr__(self, "prediction", _scalar(self.prediction, "prediction"))
|
|
71
|
+
version = _tick(self.address_version, "address_version")
|
|
72
|
+
if version < 0:
|
|
73
|
+
raise ValueError("address_version must be nonnegative")
|
|
74
|
+
object.__setattr__(self, "address_version", version)
|
|
75
|
+
if not isinstance(self.ticket_id, str) or not self.ticket_id:
|
|
76
|
+
raise ValueError("ticket_id must be a nonempty string")
|
|
77
|
+
# Immutable bytes backing prevents callers from re-enabling NumPy writes.
|
|
78
|
+
object.__setattr__(self, "read_key", _key(self.read_key).tobytes())
|
|
79
|
+
|
|
80
|
+
def __getattribute__(self, name: str) -> Any:
|
|
81
|
+
value = object.__getattribute__(self, name)
|
|
82
|
+
if name == "read_key" and isinstance(value, bytes):
|
|
83
|
+
# Return a fresh view: NumPy shape/dtype metadata remains mutable even
|
|
84
|
+
# when its data backing is immutable. Never expose our stored object.
|
|
85
|
+
return np.frombuffer(value, dtype=np.float64)
|
|
86
|
+
return value
|
|
87
|
+
|
|
88
|
+
@property
|
|
89
|
+
def available_hour(self) -> int:
|
|
90
|
+
return self.issue_hour + self.horizon
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
@dataclass(frozen=True)
|
|
94
|
+
class ReleasedEvidence:
|
|
95
|
+
ticket: ForecastTicket
|
|
96
|
+
target: float
|
|
97
|
+
now_hour: int
|
|
98
|
+
|
|
99
|
+
def __post_init__(self) -> None:
|
|
100
|
+
if not isinstance(self.ticket, ForecastTicket):
|
|
101
|
+
raise ValueError("Evidence requires a ForecastTicket")
|
|
102
|
+
now = _tick(self.now_hour, "now_hour")
|
|
103
|
+
if now < self.ticket.available_hour:
|
|
104
|
+
raise ValueError("Outcome has not matured")
|
|
105
|
+
object.__setattr__(self, "now_hour", now)
|
|
106
|
+
object.__setattr__(self, "target", _scalar(self.target, "target"))
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
@dataclass(frozen=True)
|
|
110
|
+
class BudgetReport:
|
|
111
|
+
value_bytes: int
|
|
112
|
+
metadata_bytes: int
|
|
113
|
+
pending_bytes: int
|
|
114
|
+
seconds: float
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
@runtime_checkable
|
|
118
|
+
class MemoryProtocol(Protocol):
|
|
119
|
+
def read(self, key: Any) -> Any: ...
|
|
120
|
+
|
|
121
|
+
def write(self, evidence: ReleasedEvidence) -> None: ...
|
|
122
|
+
|
|
123
|
+
def snapshot(self) -> dict[str, Any]: ...
|
revisionlab/py.typed
ADDED
|
File without changes
|
revisionlab/replay.py
ADDED
|
@@ -0,0 +1,204 @@
|
|
|
1
|
+
"""Owner-bound delayed evidence replay; pending state contains no future labels."""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
import sys
|
|
5
|
+
import uuid
|
|
6
|
+
from numbers import Integral
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import torch
|
|
11
|
+
from numpy.typing import NDArray
|
|
12
|
+
|
|
13
|
+
from .protocols import ForecastTicket, ReleasedEvidence
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class ReplayLedger:
|
|
17
|
+
"""One horizon and address version per ledger, with retained audit IDs.
|
|
18
|
+
|
|
19
|
+
Settled IDs are retained without a bound to guarantee lifetime exactly-once
|
|
20
|
+
release through checkpoints. There is no automatic retention limit: callers
|
|
21
|
+
must rotate ledgers explicitly if accepting a shorter deduplication window.
|
|
22
|
+
``budget`` exposes growing audit payload plus Python container overhead.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
def __init__(self, horizon: int, *, owner_id: str | None = None):
|
|
26
|
+
if isinstance(horizon, bool) or not isinstance(horizon, Integral) or horizon < 1:
|
|
27
|
+
raise ValueError("horizon must be a positive integer")
|
|
28
|
+
self.horizon: int = int(horizon)
|
|
29
|
+
self.owner_id: str = str(uuid.uuid4()) if owner_id is None else owner_id
|
|
30
|
+
if not isinstance(self.owner_id, str) or not self.owner_id or ":" in self.owner_id:
|
|
31
|
+
raise ValueError("owner_id must be a nonempty string without colon")
|
|
32
|
+
self.pending: dict[str, ForecastTicket] = {}
|
|
33
|
+
self.settled_ids: set[str] = set()
|
|
34
|
+
self.clock: int | None = None
|
|
35
|
+
self.squared_error: float = 0.0
|
|
36
|
+
self.base_squared_error: float = 0.0
|
|
37
|
+
self.releases: int = 0
|
|
38
|
+
|
|
39
|
+
def _time(self, hour: int) -> None:
|
|
40
|
+
if isinstance(hour, bool) or not isinstance(hour, Integral):
|
|
41
|
+
raise ValueError("Event time must be an integer hour")
|
|
42
|
+
if not -(2**63) <= hour < 2**63:
|
|
43
|
+
raise ValueError("Event time must fit int64")
|
|
44
|
+
if self.clock is not None and hour < self.clock:
|
|
45
|
+
raise ValueError("Event time must be monotonic; release before issuing at that tick")
|
|
46
|
+
|
|
47
|
+
def issue(self, ticket: ForecastTicket) -> None:
|
|
48
|
+
self._time(ticket.issue_hour)
|
|
49
|
+
if ticket.horizon != self.horizon or ticket.address_version != 0:
|
|
50
|
+
raise ValueError("Ticket horizon or address version does not match ledger")
|
|
51
|
+
if not ticket.ticket_id.startswith(self.owner_id + ":"):
|
|
52
|
+
raise ValueError("Foreign ticket owner")
|
|
53
|
+
if ticket.ticket_id in self.pending or ticket.ticket_id in self.settled_ids:
|
|
54
|
+
raise ValueError("Ticket ID has already been issued")
|
|
55
|
+
if ticket.prediction is None or not math.isfinite(ticket.prediction):
|
|
56
|
+
raise ValueError("Ledger requires the actual finite issued prediction")
|
|
57
|
+
self.pending[ticket.ticket_id] = ticket
|
|
58
|
+
self.clock = ticket.issue_hour
|
|
59
|
+
|
|
60
|
+
def validate_release(self, ticket_id: str, target: float, now_hour: int) -> ReleasedEvidence:
|
|
61
|
+
self._time(now_hour)
|
|
62
|
+
if ticket_id in self.settled_ids:
|
|
63
|
+
raise ValueError("Ticket has already been released")
|
|
64
|
+
if ticket_id not in self.pending:
|
|
65
|
+
raise ValueError("Unknown or foreign ticket ID")
|
|
66
|
+
ticket = self.pending[ticket_id]
|
|
67
|
+
if ticket.address_version != 0 or ticket.horizon != self.horizon:
|
|
68
|
+
raise ValueError("Stored ticket horizon or address version does not match ledger")
|
|
69
|
+
evidence = ReleasedEvidence(ticket, target, now_hour)
|
|
70
|
+
assert ticket.prediction is not None
|
|
71
|
+
try:
|
|
72
|
+
squared_error = (ticket.prediction - evidence.target) ** 2
|
|
73
|
+
base_error = (ticket.base_prediction - evidence.target) ** 2
|
|
74
|
+
except OverflowError as error:
|
|
75
|
+
raise ValueError("Squared error must be finite") from error
|
|
76
|
+
if not math.isfinite(
|
|
77
|
+
squared_error + self.squared_error + base_error + self.base_squared_error
|
|
78
|
+
):
|
|
79
|
+
raise ValueError("Squared error must be finite")
|
|
80
|
+
return evidence
|
|
81
|
+
|
|
82
|
+
def _settle(self, evidence: ReleasedEvidence) -> float:
|
|
83
|
+
ticket = evidence.ticket
|
|
84
|
+
canonical = self.validate_release(ticket.ticket_id, evidence.target, evidence.now_hour)
|
|
85
|
+
if canonical.ticket is not ticket:
|
|
86
|
+
raise ValueError("Release must use the canonical pending ticket")
|
|
87
|
+
assert ticket.prediction is not None
|
|
88
|
+
loss = (ticket.prediction - evidence.target) ** 2
|
|
89
|
+
self.squared_error += loss
|
|
90
|
+
self.base_squared_error += (ticket.base_prediction - evidence.target) ** 2
|
|
91
|
+
self.releases += 1
|
|
92
|
+
del self.pending[ticket.ticket_id]
|
|
93
|
+
self.settled_ids.add(ticket.ticket_id)
|
|
94
|
+
self.clock = evidence.now_hour
|
|
95
|
+
return loss
|
|
96
|
+
|
|
97
|
+
@property
|
|
98
|
+
def mse(self) -> float | None:
|
|
99
|
+
return self.squared_error / self.releases if self.releases else None
|
|
100
|
+
|
|
101
|
+
@property
|
|
102
|
+
def base_mse(self) -> float | None:
|
|
103
|
+
return self.base_squared_error / self.releases if self.releases else None
|
|
104
|
+
|
|
105
|
+
def budget(self) -> dict[str, int]:
|
|
106
|
+
return {
|
|
107
|
+
"pending_tickets": len(self.pending),
|
|
108
|
+
"settled_ids": len(self.settled_ids),
|
|
109
|
+
"pending_key_bytes": sum(ticket.read_key.nbytes for ticket in self.pending.values()),
|
|
110
|
+
"audit_id_payload_bytes": sum(len(value.encode()) for value in self.settled_ids),
|
|
111
|
+
"audit_python_bytes": sys.getsizeof(self.settled_ids)
|
|
112
|
+
+ sum(sys.getsizeof(value) for value in self.settled_ids),
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
def state(self) -> dict[str, Any]:
|
|
116
|
+
return {
|
|
117
|
+
"horizon": self.horizon,
|
|
118
|
+
"owner_id": self.owner_id,
|
|
119
|
+
"clock": self.clock,
|
|
120
|
+
"squared_error": self.squared_error,
|
|
121
|
+
"base_squared_error": self.base_squared_error,
|
|
122
|
+
"releases": self.releases,
|
|
123
|
+
"settled_ids": sorted(self.settled_ids),
|
|
124
|
+
"pending": [
|
|
125
|
+
{
|
|
126
|
+
"issue_hour": ticket.issue_hour,
|
|
127
|
+
"horizon": ticket.horizon,
|
|
128
|
+
"base_prediction": ticket.base_prediction,
|
|
129
|
+
"read_key": torch.tensor(np.array(ticket.read_key, copy=True)),
|
|
130
|
+
"prediction": ticket.prediction,
|
|
131
|
+
"ticket_id": ticket.ticket_id,
|
|
132
|
+
"address_version": ticket.address_version,
|
|
133
|
+
}
|
|
134
|
+
for ticket in self.pending.values()
|
|
135
|
+
],
|
|
136
|
+
}
|
|
137
|
+
|
|
138
|
+
@classmethod
|
|
139
|
+
def from_state(cls, state: dict[str, Any]) -> "ReplayLedger":
|
|
140
|
+
ledger = cls(state["horizon"], owner_id=state["owner_id"])
|
|
141
|
+
settled = state["settled_ids"]
|
|
142
|
+
if len(set(settled)) != len(settled) or any(
|
|
143
|
+
not isinstance(value, str) or not value.startswith(ledger.owner_id + ":")
|
|
144
|
+
for value in settled
|
|
145
|
+
):
|
|
146
|
+
raise ValueError("Invalid settled ticket IDs")
|
|
147
|
+
ledger.settled_ids = set(settled)
|
|
148
|
+
for data in sorted(state["pending"], key=lambda item: item["issue_hour"]):
|
|
149
|
+
ledger.issue(ForecastTicket(**data))
|
|
150
|
+
clock = state["clock"]
|
|
151
|
+
if clock is not None:
|
|
152
|
+
ledger._time(clock)
|
|
153
|
+
elif ledger.pending or ledger.settled_ids:
|
|
154
|
+
raise ValueError("Nonempty ledger requires a clock")
|
|
155
|
+
ledger.clock = clock
|
|
156
|
+
count = state["releases"]
|
|
157
|
+
if isinstance(count, bool) or not isinstance(count, int) or count != len(settled):
|
|
158
|
+
raise ValueError("Release count must match settled audit IDs")
|
|
159
|
+
ledger.releases = count
|
|
160
|
+
for name in ("squared_error", "base_squared_error"):
|
|
161
|
+
value = state[name]
|
|
162
|
+
if not isinstance(value, (int, float)) or not math.isfinite(value) or value < 0:
|
|
163
|
+
raise ValueError("Invalid accumulated error")
|
|
164
|
+
setattr(ledger, name, float(value))
|
|
165
|
+
return ledger
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
class DelayedReplay:
|
|
169
|
+
"""Sequential single-process replay: revise only after an issued outcome matures.
|
|
170
|
+
|
|
171
|
+
Calls must be serialized by the caller; this class is not thread-safe.
|
|
172
|
+
"""
|
|
173
|
+
|
|
174
|
+
def __init__(self, memory: Any, horizon: int, *, correction_clip: float = 3.0):
|
|
175
|
+
if not math.isfinite(correction_clip) or correction_clip <= 0:
|
|
176
|
+
raise ValueError("correction_clip must be finite and positive")
|
|
177
|
+
self.memory = memory
|
|
178
|
+
self.ledger = ReplayLedger(horizon)
|
|
179
|
+
self.correction_clip = float(correction_clip)
|
|
180
|
+
|
|
181
|
+
def issue(
|
|
182
|
+
self, issue_hour: int, base_prediction: float, read_key: NDArray[np.float64] | torch.Tensor
|
|
183
|
+
) -> ForecastTicket:
|
|
184
|
+
self.ledger._time(issue_hour)
|
|
185
|
+
correction = float(self.memory.read(read_key))
|
|
186
|
+
if not math.isfinite(correction):
|
|
187
|
+
raise ValueError("Memory correction must be finite")
|
|
188
|
+
correction = max(-self.correction_clip, min(self.correction_clip, correction))
|
|
189
|
+
ticket = ForecastTicket(
|
|
190
|
+
issue_hour,
|
|
191
|
+
self.ledger.horizon,
|
|
192
|
+
base_prediction,
|
|
193
|
+
read_key,
|
|
194
|
+
prediction=base_prediction + correction,
|
|
195
|
+
ticket_id=self.ledger.owner_id + ":" + str(uuid.uuid4()),
|
|
196
|
+
address_version=0,
|
|
197
|
+
)
|
|
198
|
+
self.ledger.issue(ticket)
|
|
199
|
+
return ticket
|
|
200
|
+
|
|
201
|
+
def release(self, ticket_id: str, target: float, now_hour: int) -> float:
|
|
202
|
+
evidence = self.ledger.validate_release(ticket_id, target, now_hour)
|
|
203
|
+
self.memory.write(evidence)
|
|
204
|
+
return self.ledger._settle(evidence)
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""Atomic, weights-only checkpoints with explicit model and ledger validation."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import tempfile
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
from torch import nn
|
|
10
|
+
|
|
11
|
+
from .baselines import NoWrite, RLS_Corrector, ShuffledWrite
|
|
12
|
+
from .memory import ResidualMemory
|
|
13
|
+
from .replay import DelayedReplay, ReplayLedger
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def save_checkpoint(path: str | Path, runner: DelayedReplay) -> None:
|
|
17
|
+
if type(runner.memory) not in (ResidualMemory, RLS_Corrector, NoWrite, ShuffledWrite):
|
|
18
|
+
raise ValueError("Checkpoint supports native RevisionLab models only")
|
|
19
|
+
if runner.memory.values.dtype not in (torch.float32, torch.float64):
|
|
20
|
+
raise ValueError("Checkpoint dtype must be float32 or float64")
|
|
21
|
+
if isinstance(runner.memory, ResidualMemory) and not 0 < runner.memory.rate <= 1:
|
|
22
|
+
raise ValueError("Checkpoint learning rate must be in (0,1]")
|
|
23
|
+
if isinstance(runner.memory, RLS_Corrector) and not 0 < runner.memory.forget <= 1:
|
|
24
|
+
raise ValueError("Checkpoint forgetting factor must be in (0,1]")
|
|
25
|
+
for name, value in runner.memory.state_dict().items():
|
|
26
|
+
if value.is_floating_point() and value.dtype != runner.memory.values.dtype:
|
|
27
|
+
raise ValueError(f"Checkpoint tensor dtype mismatch: {name}")
|
|
28
|
+
if not bool(torch.isfinite(value).all()):
|
|
29
|
+
raise ValueError(f"Checkpoint tensor is non-finite: {name}")
|
|
30
|
+
payload = {
|
|
31
|
+
"format_version": 1,
|
|
32
|
+
"model_type": type(runner.memory).__name__,
|
|
33
|
+
"model_config": runner.memory.config(),
|
|
34
|
+
"model_state": {
|
|
35
|
+
name: value.detach().cpu().clone() for name, value in runner.memory.state_dict().items()
|
|
36
|
+
},
|
|
37
|
+
"ledger": runner.ledger.state(),
|
|
38
|
+
"correction_clip": runner.correction_clip,
|
|
39
|
+
}
|
|
40
|
+
destination = Path(path)
|
|
41
|
+
destination.parent.mkdir(parents=True, exist_ok=True)
|
|
42
|
+
descriptor, temporary = tempfile.mkstemp(
|
|
43
|
+
prefix=destination.name + ".", suffix=".tmp", dir=destination.parent
|
|
44
|
+
)
|
|
45
|
+
try:
|
|
46
|
+
with os.fdopen(descriptor, "wb") as stream:
|
|
47
|
+
torch.save(payload, stream)
|
|
48
|
+
stream.flush()
|
|
49
|
+
os.fsync(stream.fileno())
|
|
50
|
+
os.replace(temporary, destination)
|
|
51
|
+
finally:
|
|
52
|
+
if os.path.exists(temporary):
|
|
53
|
+
os.unlink(temporary)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def load_checkpoint(path: str | Path) -> DelayedReplay:
|
|
57
|
+
payload: dict[str, Any] = torch.load(path, map_location="cpu", weights_only=True)
|
|
58
|
+
if not isinstance(payload, dict) or payload.get("format_version") != 1:
|
|
59
|
+
raise ValueError("Unsupported checkpoint format")
|
|
60
|
+
factories: dict[str, type[nn.Module]] = {
|
|
61
|
+
"ResidualMemory": ResidualMemory,
|
|
62
|
+
"RLS_Corrector": RLS_Corrector,
|
|
63
|
+
"NoWrite": NoWrite,
|
|
64
|
+
"ShuffledWrite": ShuffledWrite,
|
|
65
|
+
}
|
|
66
|
+
model_type = payload.get("model_type")
|
|
67
|
+
if model_type not in factories:
|
|
68
|
+
raise ValueError("Unknown checkpoint model type")
|
|
69
|
+
config = dict(payload["model_config"])
|
|
70
|
+
dtype_name = config.pop("dtype", None)
|
|
71
|
+
if dtype_name not in ("float32", "float64"):
|
|
72
|
+
raise ValueError("Checkpoint dtype must be float32 or float64")
|
|
73
|
+
dtype = torch.float32 if dtype_name == "float32" else torch.float64
|
|
74
|
+
memory = factories[model_type](**config, dtype=dtype)
|
|
75
|
+
expected = memory.state_dict()
|
|
76
|
+
saved = payload["model_state"]
|
|
77
|
+
if not isinstance(saved, dict) or set(saved) != set(expected):
|
|
78
|
+
raise ValueError("Checkpoint model state keys mismatch")
|
|
79
|
+
for name, value in saved.items():
|
|
80
|
+
if not isinstance(value, torch.Tensor):
|
|
81
|
+
raise ValueError("Checkpoint state must contain tensors")
|
|
82
|
+
if value.shape != expected[name].shape or value.dtype != expected[name].dtype:
|
|
83
|
+
raise ValueError(f"Checkpoint tensor shape/dtype mismatch: {name}")
|
|
84
|
+
if not bool(torch.isfinite(value).all()):
|
|
85
|
+
raise ValueError(f"Checkpoint tensor is non-finite: {name}")
|
|
86
|
+
if int(saved["writes"]) < 0:
|
|
87
|
+
raise ValueError("Checkpoint write counter cannot be negative")
|
|
88
|
+
if isinstance(memory, ResidualMemory):
|
|
89
|
+
if not 0 < float(saved["learning_rate"]) <= 1:
|
|
90
|
+
raise ValueError("Checkpoint learning rate must be in (0,1]")
|
|
91
|
+
if not torch.equal(saved["learning_rate"], expected["learning_rate"]):
|
|
92
|
+
raise ValueError("Checkpoint learning rate and config mismatch")
|
|
93
|
+
if isinstance(memory, RLS_Corrector):
|
|
94
|
+
if not 0 < float(saved["forgetting_factor"]) <= 1:
|
|
95
|
+
raise ValueError("Checkpoint forgetting factor must be in (0,1]")
|
|
96
|
+
if not torch.equal(saved["forgetting_factor"], expected["forgetting_factor"]):
|
|
97
|
+
raise ValueError("Checkpoint forgetting factor and config mismatch")
|
|
98
|
+
if isinstance(memory, NoWrite) and (bool(saved["values"].any()) or int(saved["writes"]) != 0):
|
|
99
|
+
raise ValueError("No-write checkpoint must preserve zero state")
|
|
100
|
+
if isinstance(memory, ShuffledWrite):
|
|
101
|
+
permutation = saved["permutation"]
|
|
102
|
+
if not torch.equal(permutation.sort().values, torch.arange(memory.values.numel())):
|
|
103
|
+
raise ValueError("Checkpoint write permutation must be bijective")
|
|
104
|
+
if not torch.equal(permutation, expected["permutation"]):
|
|
105
|
+
raise ValueError("Checkpoint write permutation and seed config mismatch")
|
|
106
|
+
memory.load_state_dict(saved, strict=True)
|
|
107
|
+
ledger = ReplayLedger.from_state(payload["ledger"])
|
|
108
|
+
for ticket in ledger.pending.values():
|
|
109
|
+
if ticket.read_key.size != memory.values.numel():
|
|
110
|
+
raise ValueError("Pending ticket address width does not match checkpoint model")
|
|
111
|
+
runner = DelayedReplay(memory, ledger.horizon, correction_clip=payload["correction_clip"])
|
|
112
|
+
runner.ledger = ledger
|
|
113
|
+
return runner
|
revisionlab/smoke.py
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""A synthetic delayed-residual smoke check, not forecasting validation."""
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import json
|
|
5
|
+
from collections import deque
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import torch
|
|
9
|
+
|
|
10
|
+
from .baselines import NoWrite, RLS_Corrector, ShuffledWrite
|
|
11
|
+
from .memory import ResidualMemory
|
|
12
|
+
from .replay import DelayedReplay
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def synthetic_check(steps: int = 128, seed: int = 0, horizon: int = 3) -> dict[str, object]:
|
|
16
|
+
if isinstance(steps, bool) or not isinstance(steps, int) or steps < 1:
|
|
17
|
+
raise ValueError("steps must be a positive integer")
|
|
18
|
+
if isinstance(horizon, bool) or not isinstance(horizon, int) or horizon < 1:
|
|
19
|
+
raise ValueError("horizon must be a positive integer")
|
|
20
|
+
generator = np.random.default_rng(seed)
|
|
21
|
+
slots = 8
|
|
22
|
+
addresses = generator.integers(slots, size=steps)
|
|
23
|
+
noise = generator.normal(0, 0.03, size=steps)
|
|
24
|
+
residuals = np.linspace(-2, 2, slots)[addresses] + noise
|
|
25
|
+
methods: dict[str, ResidualMemory | RLS_Corrector] = {
|
|
26
|
+
"residual": ResidualMemory(slots, 0.5),
|
|
27
|
+
"rls": RLS_Corrector(slots, 0.999),
|
|
28
|
+
"no-write": NoWrite(slots),
|
|
29
|
+
"shuffled-write": ShuffledWrite(slots, 0.5, seed=seed),
|
|
30
|
+
}
|
|
31
|
+
results: dict[str, object] = {}
|
|
32
|
+
for name, memory in methods.items():
|
|
33
|
+
runner = DelayedReplay(memory, horizon)
|
|
34
|
+
pending: deque[tuple[int, str, int]] = deque()
|
|
35
|
+
for now in range(steps):
|
|
36
|
+
# This evaluator queue owns targets; runner pending state contains
|
|
37
|
+
# only issued tickets and never a not-yet-mature target.
|
|
38
|
+
while pending and pending[0][0] <= now:
|
|
39
|
+
_, ticket_id, index = pending.popleft()
|
|
40
|
+
runner.release(ticket_id, float(residuals[index]), now)
|
|
41
|
+
key = np.zeros(slots)
|
|
42
|
+
key[addresses[now]] = 1
|
|
43
|
+
ticket = runner.issue(now, 0.0, key)
|
|
44
|
+
pending.append((ticket.available_hour, ticket.ticket_id, now))
|
|
45
|
+
while pending:
|
|
46
|
+
available, ticket_id, index = pending.popleft()
|
|
47
|
+
runner.release(ticket_id, float(residuals[index]), available)
|
|
48
|
+
corrected_mse = runner.ledger.mse
|
|
49
|
+
base_mse = runner.ledger.base_mse
|
|
50
|
+
assert corrected_mse is not None and base_mse is not None
|
|
51
|
+
results[name] = {
|
|
52
|
+
"base_mse": base_mse,
|
|
53
|
+
"corrected_mse": corrected_mse,
|
|
54
|
+
"mse_improvement": base_mse - corrected_mse,
|
|
55
|
+
"releases": runner.ledger.releases,
|
|
56
|
+
"model_tensor_bytes": sum(
|
|
57
|
+
value.numel() * value.element_size() for value in memory.state_dict().values()
|
|
58
|
+
),
|
|
59
|
+
"ledger": runner.ledger.budget(),
|
|
60
|
+
}
|
|
61
|
+
return {
|
|
62
|
+
"scope": "synthetic stationary address-residual smoke; not real-world forecasting evidence",
|
|
63
|
+
"steps": steps,
|
|
64
|
+
"seed": seed,
|
|
65
|
+
"horizon": horizon,
|
|
66
|
+
"dtype": str(torch.float64),
|
|
67
|
+
"methods": results,
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def main(argv: list[str] | None = None) -> None:
|
|
72
|
+
parser = argparse.ArgumentParser(prog="revisionlab-check")
|
|
73
|
+
parser.add_argument("--steps", type=int, default=128)
|
|
74
|
+
parser.add_argument("--seed", type=int, default=0)
|
|
75
|
+
parser.add_argument("--horizon", type=int, default=3)
|
|
76
|
+
parser.add_argument(
|
|
77
|
+
"--json", action="store_true", help="print machine-readable synthetic results"
|
|
78
|
+
)
|
|
79
|
+
args = parser.parse_args(argv)
|
|
80
|
+
try:
|
|
81
|
+
report = synthetic_check(args.steps, args.seed, args.horizon)
|
|
82
|
+
except (ValueError, OverflowError) as error:
|
|
83
|
+
parser.error(str(error))
|
|
84
|
+
if args.json:
|
|
85
|
+
print(json.dumps(report, indent=2, allow_nan=False))
|
|
86
|
+
else:
|
|
87
|
+
print(report["scope"])
|
|
88
|
+
methods = report["methods"]
|
|
89
|
+
assert isinstance(methods, dict)
|
|
90
|
+
for name, row in methods.items():
|
|
91
|
+
print(
|
|
92
|
+
f"{name}: base MSE={row['base_mse']:.6f}; corrected MSE={row['corrected_mse']:.6f}"
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
if __name__ == "__main__":
|
|
97
|
+
main()
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: revisionlab
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: PyTorch-native delayed-feedback memory contracts, controls, and diagnostics
|
|
5
|
+
Project-URL: Repository, https://github.com/cjw0076/revisionlab
|
|
6
|
+
Project-URL: Issues, https://github.com/cjw0076/revisionlab/issues
|
|
7
|
+
Author: cjw0076
|
|
8
|
+
License-Expression: MIT
|
|
9
|
+
License-File: LICENSE
|
|
10
|
+
Classifier: Development Status :: 3 - Alpha
|
|
11
|
+
Classifier: Programming Language :: Python :: 3
|
|
12
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
13
|
+
Requires-Python: >=3.10
|
|
14
|
+
Requires-Dist: numpy>=1.24
|
|
15
|
+
Requires-Dist: torch>=2.2
|
|
16
|
+
Provides-Extra: dev
|
|
17
|
+
Requires-Dist: build>=1.2; extra == 'dev'
|
|
18
|
+
Requires-Dist: hatch>=1.13; extra == 'dev'
|
|
19
|
+
Requires-Dist: mypy>=1.11; extra == 'dev'
|
|
20
|
+
Requires-Dist: pytest>=8; extra == 'dev'
|
|
21
|
+
Requires-Dist: ruff>=0.6; extra == 'dev'
|
|
22
|
+
Requires-Dist: twine>=5; extra == 'dev'
|
|
23
|
+
Description-Content-Type: text/markdown
|
|
24
|
+
|
|
25
|
+
# RevisionLab
|
|
26
|
+
|
|
27
|
+
**PyTorch-native delayed-feedback memory contracts, controls, and diagnostics.**
|
|
28
|
+
|
|
29
|
+
RevisionLab records a prediction when it is issued, waits for its outcome to mature, and applies a residual-state update to the memory that owned the prediction. It provides a small substrate for testing delayed correction: immutable tickets, explicit memory interfaces, replay, negative controls, and resumable state.
|
|
30
|
+
|
|
31
|
+
v0.1 focuses on fixed-capacity residual memory and RLS correction. It does not implement Cosmos v14 state splitting or the v10 learned-refinement toy. It makes no claim to support every neural architecture. [archcredit](https://github.com/cjw0076/archcredit) remains a separate architecture/credit benchmark library; RevisionLab does not depend on it.
|
|
32
|
+
|
|
33
|
+
## Install and check
|
|
34
|
+
|
|
35
|
+
From a checkout with Python 3.10 or newer:
|
|
36
|
+
|
|
37
|
+
```powershell
|
|
38
|
+
python -m pip install -e ".[dev]"
|
|
39
|
+
revisionlab-check
|
|
40
|
+
```
|
|
41
|
+
|
|
42
|
+
The check runs a small synthetic delayed-feedback example. It does not download weather data or replay the historical experiment.
|
|
43
|
+
|
|
44
|
+
## Minimal API
|
|
45
|
+
|
|
46
|
+
```python
|
|
47
|
+
import numpy as np
|
|
48
|
+
from revisionlab import ResidualMemory
|
|
49
|
+
from revisionlab.replay import DelayedReplay
|
|
50
|
+
|
|
51
|
+
runner = DelayedReplay(ResidualMemory(2, rate=0.1), horizon=2)
|
|
52
|
+
ticket = runner.issue(0, 1.0, np.array([1.0, 0.0]))
|
|
53
|
+
error = runner.release(ticket.ticket_id, target=2.0, now_hour=2)
|
|
54
|
+
assert ticket.prediction == 1.0
|
|
55
|
+
assert error == 1.0
|
|
56
|
+
assert runner.memory.read(np.array([1.0, 0.0])) == 0.1
|
|
57
|
+
```
|
|
58
|
+
|
|
59
|
+
See [examples/delayed_forecast.py](examples/delayed_forecast.py) for a complete replay example. `revisionlab.serialization.save_checkpoint(path, runner)` and `load_checkpoint(path)` preserve delayed-feedback runners using the four native built-in correctors. Custom correctors and the NumPy reference are not supported by these checkpoint helpers in v0.1.
|
|
60
|
+
|
|
61
|
+
## Contracts
|
|
62
|
+
|
|
63
|
+
| Boundary | Purpose |
|
|
64
|
+
| --- | --- |
|
|
65
|
+
| Prediction ticket | Keeps prediction-time key, baseline, actual forecast, maturity, and owner identity. |
|
|
66
|
+
| Released evidence | Rejects targets that have not matured and invalid timestamp/target fields. |
|
|
67
|
+
| Memory `read`/`write`/`snapshot` | Reads a correction, applies valid delayed evidence, and copies state. |
|
|
68
|
+
| `DelayedReplay` | Owns tickets, evaluates the issued forecast, and settles feedback exactly once. |
|
|
69
|
+
| Checkpoint | Preserves memory, replay owner, pending tickets, and settled identities for resume. |
|
|
70
|
+
| Local alignment | Compares a correction update with a local residual-loss reference direction. |
|
|
71
|
+
|
|
72
|
+
The NumPy reference and PyTorch-native implementations use fixed-capacity state. Negative controls are explicitly named; shuffled writes corrupt write addressing rather than demonstrating a new learning rule. A no-write control must preserve zero correction and perform no learning.
|
|
73
|
+
|
|
74
|
+
`ResidualMemory(slots, rate)` is the native normalized-delta corrector. `revisionlab.baselines` provides `RLS_Corrector(slots, forget)`, `NoWrite(slots)`, and `ShuffledWrite(slots, rate, seed=...)`. See [the method manifest](docs/methods.md) for equations and clipping. Low-level memory `write` applies each valid call; **exactly-once settlement is a `DelayedReplay` guarantee**, not a guarantee of direct memory writes. Tickets retain prediction-time keys and forecasts, and the runner rejects duplicate/foreign settlement.
|
|
75
|
+
|
|
76
|
+
Memory values have fixed capacity. Replay retains settled ticket identities for lifetime deduplication, so its total audit ledger grows with settled forecasts; report that overhead alongside pending tickets. Do not call the entire runner constant-memory.
|
|
77
|
+
|
|
78
|
+
## What the evidence means
|
|
79
|
+
|
|
80
|
+
Recovered v13 artifacts report the following historical mean joint MSE:
|
|
81
|
+
|
|
82
|
+
| Historical policy | Mean joint MSE |
|
|
83
|
+
| --- | ---: |
|
|
84
|
+
| Frozen baseline | 2.122014 |
|
|
85
|
+
| Native live residual | 2.092176 |
|
|
86
|
+
| Live RLS | 2.070044 |
|
|
87
|
+
|
|
88
|
+
These are prior-run receipts for one site and the first half of 2026, across five weight seeds. The residual method improved on the frozen baseline in that recorded setting, while RLS had the lower mean MSE. This release has not rerun the full weather experiment, established superiority over RLS, or demonstrated generalization to other climates. See [provenance](docs/provenance.md) and [historical evidence](docs/evidence/README.md).
|
|
89
|
+
|
|
90
|
+
The original shuffled-write policy used a fixed cyclic permutation. The release's seeded `ShuffledWrite` control is a distinct policy and does not reproduce that historical control automatically.
|
|
91
|
+
|
|
92
|
+
Local residual alignment `rho` is not full-model BPTT alignment. The delta primitive is differentiable, but that alone does not prove correct global credit assignment or causal memory capacity. Zero-norm directions have undefined alignment.
|
|
93
|
+
|
|
94
|
+
## Development and contribution
|
|
95
|
+
|
|
96
|
+
```powershell
|
|
97
|
+
python -m ruff check .
|
|
98
|
+
python -m ruff format --check .
|
|
99
|
+
python -m mypy src/revisionlab
|
|
100
|
+
python -m pytest
|
|
101
|
+
python -m build
|
|
102
|
+
python -m twine check dist/*
|
|
103
|
+
```
|
|
104
|
+
|
|
105
|
+
[CONTRIBUTING.md](CONTRIBUTING.md) explains the five-minute corrector recipe. CI covers Python 3.10/3.11/3.12, lint/format, types, tests, `aot_eager` compile/serialization/gradcheck, and wheel installation outside the checkout. A CI definition is not a claim that hosted jobs have already passed; current receipts belong in [STATE.md](STATE.md).
|
|
106
|
+
|
|
107
|
+
See [release steps](docs/releasing.md) for TestPyPI/PyPI handoff and [the source audit](docs/source-audit.md) for adversarial regression requirements. No PyPI upload automation is configured. Licensed under [MIT](LICENSE); participation follows the [Code of Conduct](CODE_OF_CONDUCT.md).
|
|
108
|
+
|
|
109
|
+
Local v0.1.0 validation: **95 tests passed, 1 CUDA hardware skip**; lint/format, strict package types, and separate reviews passed. See [verification evidence](docs/verification.md) and [release notes](docs/release-notes-v0.1.0.md). Hosted and publication receipts are recorded separately.
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
revisionlab/__init__.py,sha256=xNZ9bEB_rQF3vo0GSrduvxVV8urFLeqwbvxmRDSLmUA,541
|
|
2
|
+
revisionlab/baselines.py,sha256=5EnkAK-02lK9MIx_YKpRjkJuUkkC_h19Zxo3Z41oxB8,6520
|
|
3
|
+
revisionlab/diagnostics.py,sha256=y-QRr7Ij9FRX7TMLX60v7dGdKOU9JMg-3baAIqIEVZ8,1700
|
|
4
|
+
revisionlab/memory.py,sha256=ZB5uiBZDHPqEm_7pQiMDiKVOmpTM2vHLo-v05R9fJ0c,8525
|
|
5
|
+
revisionlab/protocols.py,sha256=WPOP35TayQhQqzRKryOkK15HznEhNT9e5ostM-ibq40,4499
|
|
6
|
+
revisionlab/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
|
|
7
|
+
revisionlab/replay.py,sha256=4U81SS9wXjm7_FO3WnrZS6Hz_9FApSjCSl99-MFGQ4w,9189
|
|
8
|
+
revisionlab/serialization.py,sha256=uhVuba79w4iVo2ba_KkhfNJ2m2YIdpFj-Nkhwfr4SII,5669
|
|
9
|
+
revisionlab/smoke.py,sha256=AuVvLpxNx3Ws5GmjcQgwgSQgFKz-ErdPLihRE3PbfDA,3880
|
|
10
|
+
revisionlab-0.1.0.dist-info/METADATA,sha256=TdzLTJV05x0O3YSts1skmKSOqWy0pm47AmYMUP2T3ts,6686
|
|
11
|
+
revisionlab-0.1.0.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
|
|
12
|
+
revisionlab-0.1.0.dist-info/entry_points.txt,sha256=CGYpMf6zgpP7cD7sNiZaxnxopitGZhp2vm_9xNDMkbo,61
|
|
13
|
+
revisionlab-0.1.0.dist-info/licenses/LICENSE,sha256=lWs6EroFZQ8AWVa1DyaiRcWRTLdYilkP8VJ1C8cPnJg,1093
|
|
14
|
+
revisionlab-0.1.0.dist-info/RECORD,,
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 cjw0076 and RevisionLab contributors
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|