interp-engine 0.0.24__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.
- interp_engine-0.0.24.dist-info/METADATA +21 -0
- interp_engine-0.0.24.dist-info/RECORD +28 -0
- interp_engine-0.0.24.dist-info/WHEEL +4 -0
- interp_engine-0.0.24.dist-info/licenses/LICENSE +19 -0
- neuron_explainer/__init__.py +0 -0
- neuron_explainer/activations/__init__.py +0 -0
- neuron_explainer/activations/activation_records.py +130 -0
- neuron_explainer/activations/activations.py +311 -0
- neuron_explainer/activations/attention_utils.py +121 -0
- neuron_explainer/activations/token_connections.py +59 -0
- neuron_explainer/api_client.py +190 -0
- neuron_explainer/azure.py +5 -0
- neuron_explainer/explanations/__init__.py +0 -0
- neuron_explainer/explanations/calibrated_simulator.py +194 -0
- neuron_explainer/explanations/explainer.py +2585 -0
- neuron_explainer/explanations/explanations.py +230 -0
- neuron_explainer/explanations/few_shot_examples.py +3125 -0
- neuron_explainer/explanations/prompt_builder.py +118 -0
- neuron_explainer/explanations/puzzles.json +399 -0
- neuron_explainer/explanations/puzzles.py +50 -0
- neuron_explainer/explanations/scoring.py +155 -0
- neuron_explainer/explanations/simulator.py +1121 -0
- neuron_explainer/explanations/test_explainer.py +227 -0
- neuron_explainer/explanations/test_simulator.py +269 -0
- neuron_explainer/explanations/token_space_few_shot_examples.py +212 -0
- neuron_explainer/fast_dataclasses/__init__.py +3 -0
- neuron_explainer/fast_dataclasses/fast_dataclasses.py +85 -0
- neuron_explainer/fast_dataclasses/test_fast_dataclasses.py +83 -0
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import List, Union
|
|
3
|
+
|
|
4
|
+
import blobfile as bf
|
|
5
|
+
from neuron_explainer.fast_dataclasses import FastDataclass, loads, register_dataclass
|
|
6
|
+
from neuron_explainer.azure import standardize_azure_url
|
|
7
|
+
import urllib.request
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@register_dataclass
|
|
11
|
+
@dataclass
|
|
12
|
+
class TokensAndWeights(FastDataclass):
|
|
13
|
+
tokens: List[str]
|
|
14
|
+
strengths: List[float]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@register_dataclass
|
|
18
|
+
@dataclass
|
|
19
|
+
class WeightBasedSummaryOfNeuron(FastDataclass):
|
|
20
|
+
input_positive: TokensAndWeights
|
|
21
|
+
input_negative: TokensAndWeights
|
|
22
|
+
output_positive: TokensAndWeights
|
|
23
|
+
output_negative: TokensAndWeights
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def load_token_weight_connections_of_neuron(
|
|
27
|
+
layer_index: Union[str, int],
|
|
28
|
+
neuron_index: Union[str, int],
|
|
29
|
+
dataset_path: str = "https://openaipublic.blob.core.windows.net/neuron-explainer/data/related-tokens/weight-based",
|
|
30
|
+
) -> WeightBasedSummaryOfNeuron:
|
|
31
|
+
"""Load the TokenLookupTableSummaryOfNeuron for the specified neuron."""
|
|
32
|
+
url = "/".join([dataset_path, str(layer_index), f"{neuron_index}.json"])
|
|
33
|
+
url = standardize_azure_url(url)
|
|
34
|
+
with urllib.request.urlopen(url) as f:
|
|
35
|
+
return loads(f.read(), backwards_compatible=False)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@register_dataclass
|
|
39
|
+
@dataclass
|
|
40
|
+
class TokenLookupTableSummaryOfNeuron(FastDataclass):
|
|
41
|
+
"""List of tokens and the average activations of a given neuron in response to each
|
|
42
|
+
respective token. These are selected from among the tokens in the vocabulary with the
|
|
43
|
+
highest average activations across an internet text dataset, with the highest activations
|
|
44
|
+
first."""
|
|
45
|
+
|
|
46
|
+
tokens: List[str]
|
|
47
|
+
average_activations: List[float]
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def load_token_lookup_table_connections_of_neuron(
|
|
51
|
+
layer_index: Union[str, int],
|
|
52
|
+
neuron_index: Union[str, int],
|
|
53
|
+
dataset_path: str = "https://openaipublic.blob.core.windows.net/neuron-explainer/data/related-tokens/activation-based",
|
|
54
|
+
) -> TokenLookupTableSummaryOfNeuron:
|
|
55
|
+
"""Load the TokenLookupTableSummaryOfNeuron for the specified neuron."""
|
|
56
|
+
url = "/".join([dataset_path, str(layer_index), f"{neuron_index}.json"])
|
|
57
|
+
url = standardize_azure_url(url)
|
|
58
|
+
with urllib.request.urlopen(url) as f:
|
|
59
|
+
return loads(f.read(), backwards_compatible=False)
|
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
import asyncio
|
|
2
|
+
import contextlib
|
|
3
|
+
import os
|
|
4
|
+
import random
|
|
5
|
+
import traceback
|
|
6
|
+
from asyncio import Semaphore
|
|
7
|
+
from functools import wraps
|
|
8
|
+
from typing import Any, Callable, Optional
|
|
9
|
+
|
|
10
|
+
import httpx
|
|
11
|
+
import orjson
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def is_api_error(err: Exception) -> bool:
|
|
15
|
+
if isinstance(err, httpx.HTTPStatusError):
|
|
16
|
+
response = err.response
|
|
17
|
+
error_data = response.json().get("error", {})
|
|
18
|
+
error_message = error_data.get("message")
|
|
19
|
+
if response.status_code in [400, 404, 415]:
|
|
20
|
+
if error_data.get("type") == "idempotency_error":
|
|
21
|
+
print(
|
|
22
|
+
f"Retrying after idempotency error: {error_message} ({response.url})"
|
|
23
|
+
)
|
|
24
|
+
return True
|
|
25
|
+
else:
|
|
26
|
+
# Invalid request
|
|
27
|
+
return False
|
|
28
|
+
else:
|
|
29
|
+
print(f"Retrying after API error: {error_message} ({response.url})")
|
|
30
|
+
return True
|
|
31
|
+
|
|
32
|
+
elif isinstance(err, httpx.ConnectError):
|
|
33
|
+
print(f"Retrying after connection error... ({err.request.url})")
|
|
34
|
+
return True
|
|
35
|
+
|
|
36
|
+
elif isinstance(err, httpx.TimeoutException):
|
|
37
|
+
print(f"Retrying after a timeout error... ({err.request.url})")
|
|
38
|
+
return True
|
|
39
|
+
|
|
40
|
+
elif isinstance(err, httpx.ReadError):
|
|
41
|
+
print(f"Retrying after a read error... ({err.request.url})")
|
|
42
|
+
return True
|
|
43
|
+
|
|
44
|
+
print(f"Retrying after an unexpected error: {repr(err)}")
|
|
45
|
+
traceback.print_tb(err.__traceback__)
|
|
46
|
+
return True
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def exponential_backoff(
|
|
50
|
+
retry_on: Callable[[Exception], bool] = lambda err: True,
|
|
51
|
+
) -> Callable[[Callable], Callable]:
|
|
52
|
+
"""
|
|
53
|
+
Returns a decorator which retries the wrapped function as long as the specified retry_on
|
|
54
|
+
function returns True for the exception, applying exponential backoff with jitter after
|
|
55
|
+
failures, up to a retry limit.
|
|
56
|
+
"""
|
|
57
|
+
init_delay_s = 1.0
|
|
58
|
+
max_delay_s = 10.0
|
|
59
|
+
# Roughly 30 minutes before we give up.
|
|
60
|
+
max_tries = 200
|
|
61
|
+
backoff_multiplier = 2.0
|
|
62
|
+
jitter = 0.2
|
|
63
|
+
|
|
64
|
+
def decorate(f: Callable) -> Callable:
|
|
65
|
+
assert asyncio.iscoroutinefunction(f)
|
|
66
|
+
|
|
67
|
+
@wraps(f)
|
|
68
|
+
async def f_retry(*args: Any, **kwargs: Any) -> None:
|
|
69
|
+
delay_s = init_delay_s
|
|
70
|
+
for i in range(max_tries):
|
|
71
|
+
try:
|
|
72
|
+
return await f(*args, **kwargs)
|
|
73
|
+
except Exception as err:
|
|
74
|
+
if not retry_on(err) or i == max_tries - 1:
|
|
75
|
+
raise
|
|
76
|
+
jittered_delay = random.uniform(
|
|
77
|
+
delay_s * (1 - jitter), delay_s * (1 + jitter)
|
|
78
|
+
)
|
|
79
|
+
await asyncio.sleep(jittered_delay)
|
|
80
|
+
delay_s = min(delay_s * backoff_multiplier, max_delay_s)
|
|
81
|
+
|
|
82
|
+
return f_retry
|
|
83
|
+
|
|
84
|
+
return decorate
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class ApiClient:
|
|
88
|
+
"""Performs inference using the OpenAI API. Supports response caching and concurrency limits."""
|
|
89
|
+
|
|
90
|
+
BASE_API_URL = "https://api.openai.com/v1"
|
|
91
|
+
|
|
92
|
+
def __init__(
|
|
93
|
+
self,
|
|
94
|
+
model_name: str,
|
|
95
|
+
# If set, no more than this number of HTTP requests will be made concurrently.
|
|
96
|
+
max_concurrent: Optional[int] = None,
|
|
97
|
+
# Whether to cache request/response pairs in memory to avoid duplicating requests.
|
|
98
|
+
cache: bool = False,
|
|
99
|
+
base_api_url: str = BASE_API_URL,
|
|
100
|
+
override_api_key: str | None = None,
|
|
101
|
+
):
|
|
102
|
+
self.model_name = model_name
|
|
103
|
+
self.base_api_url = base_api_url
|
|
104
|
+
self.override_api_key = override_api_key
|
|
105
|
+
if max_concurrent is not None:
|
|
106
|
+
self._concurrency_check: Optional[Semaphore] = Semaphore(max_concurrent)
|
|
107
|
+
else:
|
|
108
|
+
self._concurrency_check = None
|
|
109
|
+
|
|
110
|
+
if cache:
|
|
111
|
+
self._cache: Optional[dict[str, Any]] = {}
|
|
112
|
+
else:
|
|
113
|
+
self._cache = None
|
|
114
|
+
|
|
115
|
+
@exponential_backoff(retry_on=is_api_error)
|
|
116
|
+
async def make_request(
|
|
117
|
+
self,
|
|
118
|
+
timeout_seconds: Optional[int] = None,
|
|
119
|
+
json_mode: Optional[bool] = False,
|
|
120
|
+
**kwargs: Any,
|
|
121
|
+
) -> dict[str, Any]:
|
|
122
|
+
api_http_headers = {
|
|
123
|
+
"Content-Type": "application/json",
|
|
124
|
+
"Authorization": f"Bearer {os.getenv('OPENAI_API_KEY') if self.override_api_key is None else self.override_api_key}",
|
|
125
|
+
}
|
|
126
|
+
if self._cache is not None:
|
|
127
|
+
key = orjson.dumps(kwargs)
|
|
128
|
+
if key in self._cache:
|
|
129
|
+
return self._cache[key]
|
|
130
|
+
async with contextlib.AsyncExitStack() as stack:
|
|
131
|
+
if self._concurrency_check is not None:
|
|
132
|
+
await stack.enter_async_context(self._concurrency_check)
|
|
133
|
+
http_client = await stack.enter_async_context(
|
|
134
|
+
httpx.AsyncClient(timeout=timeout_seconds)
|
|
135
|
+
)
|
|
136
|
+
# If the request has a "messages" key, it should be sent to the /chat/completions
|
|
137
|
+
# endpoint. Otherwise, it should be sent to the /completions endpoint.
|
|
138
|
+
url = self.base_api_url + (
|
|
139
|
+
"/chat/completions" if "messages" in kwargs else "/completions"
|
|
140
|
+
)
|
|
141
|
+
kwargs["model"] = self.model_name
|
|
142
|
+
if json_mode:
|
|
143
|
+
kwargs["response_format"] = {"type": "json_object"}
|
|
144
|
+
# Convert max_tokens to max_completion_tokens for newer OpenAI models
|
|
145
|
+
# that don't support the legacy max_tokens parameter
|
|
146
|
+
if "max_tokens" in kwargs and "messages" in kwargs:
|
|
147
|
+
kwargs["max_completion_tokens"] = kwargs.pop("max_tokens")
|
|
148
|
+
response = await http_client.post(
|
|
149
|
+
url, headers=api_http_headers, json=kwargs
|
|
150
|
+
)
|
|
151
|
+
# # Print token usage information if available in the response
|
|
152
|
+
# try:
|
|
153
|
+
# response_data = response.json()
|
|
154
|
+
# if "usage" in response_data:
|
|
155
|
+
# usage = response_data["usage"]
|
|
156
|
+
# print(f"Token usage - Prompt: {usage.get('prompt_tokens', 'N/A')}, "
|
|
157
|
+
# f"Completion: {usage.get('completion_tokens', 'N/A')}, "
|
|
158
|
+
# f"Total: {usage.get('total_tokens', 'N/A')}")
|
|
159
|
+
# except Exception:
|
|
160
|
+
# pass # Silently ignore if we can't parse the response or find usage info
|
|
161
|
+
# The response json has useful information but the exception doesn't include it, so print it
|
|
162
|
+
# out then reraise.
|
|
163
|
+
try:
|
|
164
|
+
response.raise_for_status()
|
|
165
|
+
except Exception as e:
|
|
166
|
+
try:
|
|
167
|
+
print(f"Error response status code: {response.status_code}")
|
|
168
|
+
print(f"Error response JSON: {response.json()}")
|
|
169
|
+
except Exception:
|
|
170
|
+
print("Could not parse error response as JSON")
|
|
171
|
+
print(f"Error response text: {response.text}")
|
|
172
|
+
raise e
|
|
173
|
+
if self._cache is not None:
|
|
174
|
+
self._cache[key] = response.json()
|
|
175
|
+
response_json = response.json()
|
|
176
|
+
# print(f"response_json: {response_json}")
|
|
177
|
+
return response_json
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
if __name__ == "__main__":
|
|
181
|
+
|
|
182
|
+
async def main() -> None:
|
|
183
|
+
client = ApiClient(model_name="gpt-3.5-turbo", max_concurrent=1)
|
|
184
|
+
print(
|
|
185
|
+
await client.make_request(
|
|
186
|
+
prompt="Why did the chicken cross the road?", max_tokens=9
|
|
187
|
+
)
|
|
188
|
+
)
|
|
189
|
+
|
|
190
|
+
asyncio.run(main())
|
|
File without changes
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Code for calibrating simulations of neuron behavior. Calibration refers to a process of mapping from
|
|
3
|
+
a space of predicted activation values (e.g. [0, 10]) to the real activation distribution for a
|
|
4
|
+
neuron.
|
|
5
|
+
|
|
6
|
+
See http://go/neuron_explanation_methodology for description of calibration step. Necessary for
|
|
7
|
+
simulating neurons in the context of ablate-to-simulation, but can be skipped when using correlation
|
|
8
|
+
scoring. (Calibration may still improve quality for scoring, at least for non-linear calibration
|
|
9
|
+
methods.)
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
import asyncio
|
|
15
|
+
from abc import abstractmethod
|
|
16
|
+
from typing import Optional, Sequence
|
|
17
|
+
|
|
18
|
+
import numpy as np
|
|
19
|
+
from neuron_explainer.activations.activations import ActivationRecord
|
|
20
|
+
from neuron_explainer.explanations.explanations import ActivationScale
|
|
21
|
+
from neuron_explainer.explanations.simulator import NeuronSimulator, SequenceSimulation
|
|
22
|
+
from sklearn import linear_model
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class CalibratedNeuronSimulator(NeuronSimulator):
|
|
26
|
+
"""
|
|
27
|
+
Wrap a NeuronSimulator and calibrate it to map from the predicted activation space to the
|
|
28
|
+
actual neuron activation space.
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(self, uncalibrated_simulator: NeuronSimulator):
|
|
32
|
+
self.uncalibrated_simulator = uncalibrated_simulator
|
|
33
|
+
|
|
34
|
+
@classmethod
|
|
35
|
+
async def create(
|
|
36
|
+
cls,
|
|
37
|
+
uncalibrated_simulator: NeuronSimulator,
|
|
38
|
+
calibration_activation_records: Sequence[ActivationRecord],
|
|
39
|
+
) -> CalibratedNeuronSimulator:
|
|
40
|
+
"""
|
|
41
|
+
Create and calibrate a calibrated simulator (so initialization and calibration can be done
|
|
42
|
+
in one call).
|
|
43
|
+
"""
|
|
44
|
+
calibrated_simulator = cls(uncalibrated_simulator)
|
|
45
|
+
await calibrated_simulator.calibrate(calibration_activation_records)
|
|
46
|
+
return calibrated_simulator
|
|
47
|
+
|
|
48
|
+
async def calibrate(self, calibration_activation_records: Sequence[ActivationRecord]) -> None:
|
|
49
|
+
"""
|
|
50
|
+
Determine parameters to map from the predicted activation space to the real neuron
|
|
51
|
+
activation space, based on a calibration set.
|
|
52
|
+
|
|
53
|
+
Use when simulated sequences haven't already been produced on the calibration set.
|
|
54
|
+
"""
|
|
55
|
+
simulations = await asyncio.gather(
|
|
56
|
+
*[
|
|
57
|
+
self.uncalibrated_simulator.simulate(activations.tokens)
|
|
58
|
+
for activations in calibration_activation_records
|
|
59
|
+
]
|
|
60
|
+
)
|
|
61
|
+
self.calibrate_from_simulations(calibration_activation_records, simulations)
|
|
62
|
+
|
|
63
|
+
def calibrate_from_simulations(
|
|
64
|
+
self,
|
|
65
|
+
calibration_activation_records: Sequence[ActivationRecord],
|
|
66
|
+
simulations: Sequence[SequenceSimulation],
|
|
67
|
+
) -> None:
|
|
68
|
+
"""
|
|
69
|
+
Determine parameters to map from the predicted activation space to the real neuron
|
|
70
|
+
activation space, based on a calibration set.
|
|
71
|
+
|
|
72
|
+
Use when simulated sequences have already been produced on the calibration set.
|
|
73
|
+
"""
|
|
74
|
+
flattened_activations = []
|
|
75
|
+
flattened_simulated_activations: list[float] = []
|
|
76
|
+
for activations, simulation in zip(calibration_activation_records, simulations):
|
|
77
|
+
flattened_activations.extend(activations.activations)
|
|
78
|
+
flattened_simulated_activations.extend(simulation.expected_activations)
|
|
79
|
+
self._calibrate_from_flattened_activations(
|
|
80
|
+
np.array(flattened_activations), np.array(flattened_simulated_activations)
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
@abstractmethod
|
|
84
|
+
def _calibrate_from_flattened_activations(
|
|
85
|
+
self,
|
|
86
|
+
true_activations: np.ndarray,
|
|
87
|
+
uncalibrated_activations: np.ndarray,
|
|
88
|
+
) -> None:
|
|
89
|
+
"""
|
|
90
|
+
Determine parameters to map from the predicted activation space to the real neuron
|
|
91
|
+
activation space, based on a calibration set.
|
|
92
|
+
|
|
93
|
+
Take numpy arrays of all true activations and all uncalibrated activations on the
|
|
94
|
+
calibration set over all sequences.
|
|
95
|
+
"""
|
|
96
|
+
|
|
97
|
+
@abstractmethod
|
|
98
|
+
def apply_calibration(self, values: Sequence[float]) -> list[float]:
|
|
99
|
+
"""Apply the learned calibration to a sequence of values."""
|
|
100
|
+
|
|
101
|
+
async def simulate(self, tokens: Sequence[str]) -> SequenceSimulation:
|
|
102
|
+
uncalibrated_seq_simulation = await self.uncalibrated_simulator.simulate(tokens)
|
|
103
|
+
calibrated_activations = self.apply_calibration(
|
|
104
|
+
uncalibrated_seq_simulation.expected_activations
|
|
105
|
+
)
|
|
106
|
+
calibrated_distribution_values = [
|
|
107
|
+
self.apply_calibration(dv) for dv in uncalibrated_seq_simulation.distribution_values
|
|
108
|
+
]
|
|
109
|
+
return SequenceSimulation(
|
|
110
|
+
tokens=uncalibrated_seq_simulation.tokens,
|
|
111
|
+
expected_activations=calibrated_activations,
|
|
112
|
+
activation_scale=ActivationScale.NEURON_ACTIVATIONS,
|
|
113
|
+
distribution_values=calibrated_distribution_values,
|
|
114
|
+
distribution_probabilities=uncalibrated_seq_simulation.distribution_probabilities,
|
|
115
|
+
uncalibrated_simulation=uncalibrated_seq_simulation,
|
|
116
|
+
)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
class UncalibratedNeuronSimulator(CalibratedNeuronSimulator):
|
|
120
|
+
"""Pass through the activations without trying to calibrate."""
|
|
121
|
+
|
|
122
|
+
def __init__(self, uncalibrated_simulator: NeuronSimulator):
|
|
123
|
+
super().__init__(uncalibrated_simulator)
|
|
124
|
+
|
|
125
|
+
async def calibrate(self, calibration_activation_records: Sequence[ActivationRecord]) -> None:
|
|
126
|
+
pass
|
|
127
|
+
|
|
128
|
+
def _calibrate_from_flattened_activations(
|
|
129
|
+
self,
|
|
130
|
+
true_activations: np.ndarray,
|
|
131
|
+
uncalibrated_activations: np.ndarray,
|
|
132
|
+
) -> None:
|
|
133
|
+
pass
|
|
134
|
+
|
|
135
|
+
def apply_calibration(self, values: Sequence[float]) -> list[float]:
|
|
136
|
+
return values if isinstance(values, list) else list(values)
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
class LinearCalibratedNeuronSimulator(CalibratedNeuronSimulator):
|
|
140
|
+
"""Find a linear mapping from uncalibrated activations to true activations.
|
|
141
|
+
|
|
142
|
+
Should not change ev_correlation_score because it is invariant to linear transformations.
|
|
143
|
+
"""
|
|
144
|
+
|
|
145
|
+
def __init__(self, uncalibrated_simulator: NeuronSimulator):
|
|
146
|
+
super().__init__(uncalibrated_simulator)
|
|
147
|
+
self._regression: Optional[linear_model.LinearRegression] = None
|
|
148
|
+
|
|
149
|
+
def _calibrate_from_flattened_activations(
|
|
150
|
+
self,
|
|
151
|
+
true_activations: np.ndarray,
|
|
152
|
+
uncalibrated_activations: np.ndarray,
|
|
153
|
+
) -> None:
|
|
154
|
+
self._regression = linear_model.LinearRegression()
|
|
155
|
+
self._regression.fit(uncalibrated_activations.reshape(-1, 1), true_activations)
|
|
156
|
+
|
|
157
|
+
def apply_calibration(self, values: Sequence[float]) -> list[float]:
|
|
158
|
+
if self._regression is None:
|
|
159
|
+
raise ValueError("Must call calibrate() before apply_calibration")
|
|
160
|
+
if len(values) == 0:
|
|
161
|
+
return []
|
|
162
|
+
return self._regression.predict(np.reshape(np.array(values), (-1, 1))).tolist()
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
class PercentileMatchingCalibratedNeuronSimulator(CalibratedNeuronSimulator):
|
|
166
|
+
"""
|
|
167
|
+
Map the nth percentile of the uncalibrated activations to the nth percentile of the true
|
|
168
|
+
activations for all n.
|
|
169
|
+
|
|
170
|
+
This will match the distribution of true activations on the calibration set, but will be
|
|
171
|
+
overconfident outside of the calibration set.
|
|
172
|
+
"""
|
|
173
|
+
|
|
174
|
+
def __init__(self, uncalibrated_simulator: NeuronSimulator):
|
|
175
|
+
super().__init__(uncalibrated_simulator)
|
|
176
|
+
self._uncalibrated_activations: Optional[np.ndarray] = None
|
|
177
|
+
self._true_activations: Optional[np.ndarray] = None
|
|
178
|
+
|
|
179
|
+
def _calibrate_from_flattened_activations(
|
|
180
|
+
self,
|
|
181
|
+
true_activations: np.ndarray,
|
|
182
|
+
uncalibrated_activations: np.ndarray,
|
|
183
|
+
) -> None:
|
|
184
|
+
self._uncalibrated_activations = np.sort(uncalibrated_activations)
|
|
185
|
+
self._true_activations = np.sort(true_activations)
|
|
186
|
+
|
|
187
|
+
def apply_calibration(self, values: Sequence[float]) -> list[float]:
|
|
188
|
+
if self._true_activations is None or self._uncalibrated_activations is None:
|
|
189
|
+
raise ValueError("Must call calibrate() before apply_calibration")
|
|
190
|
+
if len(values) == 0:
|
|
191
|
+
return []
|
|
192
|
+
return np.interp(
|
|
193
|
+
np.array(values), self._uncalibrated_activations, self._true_activations
|
|
194
|
+
).tolist()
|