Skip to content

Commit f2ebe17

Browse files
committed
Deduplicate MLX activation observer adapters
1 parent a08d5af commit f2ebe17

8 files changed

Lines changed: 21 additions & 311 deletions

File tree

Cargo.toml

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,6 @@ eredu-evaluation = { version = "0.1.0", path = "eredu-evaluation" }
5151
thiserror = "2"
5252
serde = { version = "1", features = ["derive"] }
5353
serde_json = "1"
54-
tokenizers = "0.23"
5554
clap = { version = "4", features = ["derive"] }
5655
anyhow = "1"
5756
syn = { version = "3", features = ["full"] }
@@ -60,7 +59,6 @@ proc-macro2 = "1"
6059
itertools = "0.15"
6160
bindgen = "0.72"
6261
cmake = "0.1"
63-
cc = "1"
6462
tempfile = "3"
6563
half = "2"
6664
sha2 = "0.11"

eredu-backend-mlx/src/composition/deepseek.rs

Lines changed: 1 addition & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -110,54 +110,6 @@ fn construct_v4_unit(
110110
.map_err(neutral_error)
111111
}
112112

113-
struct NeutralDeepSeekObserver<'a> {
114-
inner: &'a mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
115-
}
116-
117-
impl eredu_runtime::ActivationObserver<crate::MlxTensor, eredu_nn::Error>
118-
for NeutralDeepSeekObserver<'_>
119-
{
120-
fn observe(&mut self, path: &str, value: &crate::MlxTensor) -> Result<(), eredu_nn::Error> {
121-
self.inner
122-
.observe(path, value.as_array())
123-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
124-
}
125-
126-
fn intervene(
127-
&mut self,
128-
path: &str,
129-
value: &crate::MlxTensor,
130-
) -> Result<Option<crate::MlxTensor>, eredu_nn::Error> {
131-
self.inner
132-
.intervene(path, value.as_array())
133-
.map(|value| value.map(crate::MlxTensor::from_array))
134-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
135-
}
136-
137-
fn observe_routing(
138-
&mut self,
139-
routing: eredu_runtime::RoutingObservation<'_, crate::MlxTensor>,
140-
) -> Result<(), eredu_nn::Error> {
141-
let raw = eredu_runtime::RoutingObservation {
142-
path: routing.path,
143-
selected_experts: routing.selected_experts.as_array(),
144-
selected_scores: routing.selected_scores.as_array(),
145-
route_weights: routing.route_weights.as_array(),
146-
routed_output: routing.routed_output.as_array(),
147-
local_routed_output: routing.local_routed_output.map(crate::MlxTensor::as_array),
148-
reduced_routed_output: routing
149-
.reduced_routed_output
150-
.map(crate::MlxTensor::as_array),
151-
shared_output: routing.shared_output.map(crate::MlxTensor::as_array),
152-
combined_output: routing.combined_output.map(crate::MlxTensor::as_array),
153-
expert_count: routing.expert_count,
154-
};
155-
self.inner
156-
.observe_routing(raw)
157-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
158-
}
159-
}
160-
161113
fn neutral_embedded_input<'a>(
162114
input: deepseek::mtp::EmbeddedInput<'a, Array>,
163115
) -> deepseek::mtp::EmbeddedInput<'a, crate::MlxTensor> {
@@ -1361,7 +1313,7 @@ impl DeepSeekModel {
13611313
} else {
13621314
eredu_runtime::ExpertPass::Decode
13631315
};
1364-
let mut observer = NeutralDeepSeekObserver { inner: observer };
1316+
let mut observer = crate::composition::NeutralActivationObserver::new(observer);
13651317
let output = match (&mut self.inner, &mut state.inner) {
13661318
(
13671319
DeepSeekModelInner::V3 {

eredu-backend-mlx/src/composition/gpt_oss.rs

Lines changed: 2 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -87,54 +87,6 @@ fn expert_parameter_targets(
8787
Ok(targets)
8888
}
8989

90-
struct NeutralGptOssObserver<'a> {
91-
inner: &'a mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
92-
}
93-
94-
impl eredu_runtime::ActivationObserver<crate::MlxTensor, eredu_nn::Error>
95-
for NeutralGptOssObserver<'_>
96-
{
97-
fn observe(&mut self, path: &str, value: &crate::MlxTensor) -> Result<(), eredu_nn::Error> {
98-
self.inner
99-
.observe(path, value.as_array())
100-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
101-
}
102-
103-
fn intervene(
104-
&mut self,
105-
path: &str,
106-
value: &crate::MlxTensor,
107-
) -> Result<Option<crate::MlxTensor>, eredu_nn::Error> {
108-
self.inner
109-
.intervene(path, value.as_array())
110-
.map(|value| value.map(crate::MlxTensor::from_array))
111-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
112-
}
113-
114-
fn observe_routing(
115-
&mut self,
116-
routing: eredu_runtime::RoutingObservation<'_, crate::MlxTensor>,
117-
) -> Result<(), eredu_nn::Error> {
118-
let routing = eredu_runtime::RoutingObservation {
119-
path: routing.path,
120-
selected_experts: routing.selected_experts.as_array(),
121-
selected_scores: routing.selected_scores.as_array(),
122-
route_weights: routing.route_weights.as_array(),
123-
routed_output: routing.routed_output.as_array(),
124-
local_routed_output: routing.local_routed_output.map(crate::MlxTensor::as_array),
125-
reduced_routed_output: routing
126-
.reduced_routed_output
127-
.map(crate::MlxTensor::as_array),
128-
shared_output: routing.shared_output.map(crate::MlxTensor::as_array),
129-
combined_output: routing.combined_output.map(crate::MlxTensor::as_array),
130-
expert_count: routing.expert_count,
131-
};
132-
self.inner
133-
.observe_routing(routing)
134-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
135-
}
136-
}
137-
13890
fn require_decoder_group(architecture: &NeutralArchitecture, group: usize) -> Result<(), Error> {
13991
let transport = <NeutralArchitecture as eredu_runtime::LayeredArchitecture<
14092
MlxNeuralBackend,
@@ -1317,7 +1269,7 @@ impl GptOssModel {
13171269
observer: &mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
13181270
) -> Result<Array, Error> {
13191271
let expert_cache = self.expert_cache.take();
1320-
let mut observer = NeutralGptOssObserver { inner: observer };
1272+
let mut observer = crate::composition::NeutralActivationObserver::new(observer);
13211273
let result = match expert_cache.as_ref() {
13221274
Some(expert_cache) => {
13231275
let args = self.args.clone();
@@ -1354,7 +1306,7 @@ impl GptOssModel {
13541306
cache: &mut Cache,
13551307
provider: &mut P,
13561308
stream: &Stream,
1357-
observer: &mut NeutralGptOssObserver<'_>,
1309+
observer: &mut crate::composition::NeutralActivationObserver<'_>,
13581310
) -> Result<Array, Error>
13591311
where
13601312
P: eredu_runtime::RoutedExpertProvider<MlxNeuralBackend>,

eredu-backend-mlx/src/composition/kimi_linear.rs

Lines changed: 2 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -112,54 +112,6 @@ impl KimiLinearCheckpointTemplate {
112112
}
113113
}
114114

115-
struct NeutralKimiLinearObserver<'a> {
116-
inner: &'a mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
117-
}
118-
119-
impl eredu_runtime::ActivationObserver<crate::MlxTensor, eredu_nn::Error>
120-
for NeutralKimiLinearObserver<'_>
121-
{
122-
fn observe(&mut self, path: &str, value: &crate::MlxTensor) -> Result<(), eredu_nn::Error> {
123-
self.inner
124-
.observe(path, value.as_array())
125-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
126-
}
127-
128-
fn intervene(
129-
&mut self,
130-
path: &str,
131-
value: &crate::MlxTensor,
132-
) -> Result<Option<crate::MlxTensor>, eredu_nn::Error> {
133-
self.inner
134-
.intervene(path, value.as_array())
135-
.map(|value| value.map(crate::MlxTensor::from_array))
136-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
137-
}
138-
139-
fn observe_routing(
140-
&mut self,
141-
routing: eredu_runtime::RoutingObservation<'_, crate::MlxTensor>,
142-
) -> Result<(), eredu_nn::Error> {
143-
let routing = eredu_runtime::RoutingObservation {
144-
path: routing.path,
145-
selected_experts: routing.selected_experts.as_array(),
146-
selected_scores: routing.selected_scores.as_array(),
147-
route_weights: routing.route_weights.as_array(),
148-
routed_output: routing.routed_output.as_array(),
149-
local_routed_output: routing.local_routed_output.map(crate::MlxTensor::as_array),
150-
reduced_routed_output: routing
151-
.reduced_routed_output
152-
.map(crate::MlxTensor::as_array),
153-
shared_output: routing.shared_output.map(crate::MlxTensor::as_array),
154-
combined_output: routing.combined_output.map(crate::MlxTensor::as_array),
155-
expert_count: routing.expert_count,
156-
};
157-
self.inner
158-
.observe_routing(routing)
159-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
160-
}
161-
}
162-
163115
#[derive(Clone)]
164116
struct KimiLinearUnitPopulator {
165117
external_experts: bool,
@@ -978,7 +930,7 @@ impl KimiLinearModel {
978930
) -> Result<Array, Error> {
979931
let expert_cache = self.expert_cache.take();
980932
let result = {
981-
let mut observer = NeutralKimiLinearObserver { inner: observer };
933+
let mut observer = crate::composition::NeutralActivationObserver::new(observer);
982934
match expert_cache.as_ref() {
983935
Some(expert_cache) => {
984936
let args = self.args.clone();
@@ -1016,7 +968,7 @@ impl KimiLinearModel {
1016968
cache: &mut MlxHybridState,
1017969
provider: &mut P,
1018970
stream: &Stream,
1019-
observer: &mut NeutralKimiLinearObserver<'_>,
971+
observer: &mut crate::composition::NeutralActivationObserver<'_>,
1020972
) -> Result<Array, Error>
1021973
where
1022974
P: eredu_runtime::RoutedExpertProvider<MlxNeuralBackend>,

eredu-backend-mlx/src/composition/lfm2.rs

Lines changed: 2 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -130,54 +130,6 @@ impl Lfm2CheckpointTemplate {
130130
}
131131
}
132132

133-
struct NeutralLfm2Observer<'a> {
134-
inner: &'a mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
135-
}
136-
137-
impl eredu_runtime::ActivationObserver<crate::MlxTensor, eredu_nn::Error>
138-
for NeutralLfm2Observer<'_>
139-
{
140-
fn observe(&mut self, path: &str, value: &crate::MlxTensor) -> Result<(), eredu_nn::Error> {
141-
self.inner
142-
.observe(path, value.as_array())
143-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
144-
}
145-
146-
fn intervene(
147-
&mut self,
148-
path: &str,
149-
value: &crate::MlxTensor,
150-
) -> Result<Option<crate::MlxTensor>, eredu_nn::Error> {
151-
self.inner
152-
.intervene(path, value.as_array())
153-
.map(|value| value.map(crate::MlxTensor::from_array))
154-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
155-
}
156-
157-
fn observe_routing(
158-
&mut self,
159-
routing: eredu_runtime::RoutingObservation<'_, crate::MlxTensor>,
160-
) -> Result<(), eredu_nn::Error> {
161-
let routing = eredu_runtime::RoutingObservation {
162-
path: routing.path,
163-
selected_experts: routing.selected_experts.as_array(),
164-
selected_scores: routing.selected_scores.as_array(),
165-
route_weights: routing.route_weights.as_array(),
166-
routed_output: routing.routed_output.as_array(),
167-
local_routed_output: routing.local_routed_output.map(crate::MlxTensor::as_array),
168-
reduced_routed_output: routing
169-
.reduced_routed_output
170-
.map(crate::MlxTensor::as_array),
171-
shared_output: routing.shared_output.map(crate::MlxTensor::as_array),
172-
combined_output: routing.combined_output.map(crate::MlxTensor::as_array),
173-
expert_count: routing.expert_count,
174-
};
175-
self.inner
176-
.observe_routing(routing)
177-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
178-
}
179-
}
180-
181133
#[derive(Clone)]
182134
struct Lfm2UnitPopulator {
183135
external_experts: bool,
@@ -997,7 +949,7 @@ impl Lfm2Model {
997949
) -> Result<Array, Error> {
998950
let expert_cache = self.expert_cache.take();
999951
let result = {
1000-
let mut observer = NeutralLfm2Observer { inner: observer };
952+
let mut observer = crate::composition::NeutralActivationObserver::new(observer);
1001953
match expert_cache.as_ref() {
1002954
Some(expert_cache) => {
1003955
let args = self.args.clone();
@@ -1035,7 +987,7 @@ impl Lfm2Model {
1035987
cache: &mut MlxHybridState,
1036988
provider: &mut P,
1037989
stream: &Stream,
1038-
observer: &mut NeutralLfm2Observer<'_>,
990+
observer: &mut crate::composition::NeutralActivationObserver<'_>,
1039991
) -> Result<Array, Error>
1040992
where
1041993
P: eredu_runtime::RoutedExpertProvider<MlxNeuralBackend>,

eredu-backend-mlx/src/composition/nemotron_h.rs

Lines changed: 2 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -112,54 +112,6 @@ impl NemotronHCheckpointTemplate {
112112
}
113113
}
114114

115-
struct NeutralNemotronHObserver<'a> {
116-
inner: &'a mut dyn eredu_runtime::ActivationObserver<Array, Exception>,
117-
}
118-
119-
impl eredu_runtime::ActivationObserver<crate::MlxTensor, eredu_nn::Error>
120-
for NeutralNemotronHObserver<'_>
121-
{
122-
fn observe(&mut self, path: &str, value: &crate::MlxTensor) -> Result<(), eredu_nn::Error> {
123-
self.inner
124-
.observe(path, value.as_array())
125-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
126-
}
127-
128-
fn intervene(
129-
&mut self,
130-
path: &str,
131-
value: &crate::MlxTensor,
132-
) -> Result<Option<crate::MlxTensor>, eredu_nn::Error> {
133-
self.inner
134-
.intervene(path, value.as_array())
135-
.map(|value| value.map(crate::MlxTensor::from_array))
136-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
137-
}
138-
139-
fn observe_routing(
140-
&mut self,
141-
routing: eredu_runtime::RoutingObservation<'_, crate::MlxTensor>,
142-
) -> Result<(), eredu_nn::Error> {
143-
let routing = eredu_runtime::RoutingObservation {
144-
path: routing.path,
145-
selected_experts: routing.selected_experts.as_array(),
146-
selected_scores: routing.selected_scores.as_array(),
147-
route_weights: routing.route_weights.as_array(),
148-
routed_output: routing.routed_output.as_array(),
149-
local_routed_output: routing.local_routed_output.map(crate::MlxTensor::as_array),
150-
reduced_routed_output: routing
151-
.reduced_routed_output
152-
.map(crate::MlxTensor::as_array),
153-
shared_output: routing.shared_output.map(crate::MlxTensor::as_array),
154-
combined_output: routing.combined_output.map(crate::MlxTensor::as_array),
155-
expert_count: routing.expert_count,
156-
};
157-
self.inner
158-
.observe_routing(routing)
159-
.map_err(|error| eredu_nn::Error::backend(error.to_string()))
160-
}
161-
}
162-
163115
fn neutral_embedded_input<'a>(
164116
input: eredu_architectures::nemotron_h::EmbeddedInput<'a, Array>,
165117
) -> eredu_architectures::nemotron_h::EmbeddedInput<'a, crate::MlxTensor> {
@@ -1367,7 +1319,7 @@ impl NemotronHModel {
13671319
) -> Result<Array, Error> {
13681320
let expert_cache = self.expert_cache.take();
13691321
let result = {
1370-
let mut observer = NeutralNemotronHObserver { inner: observer };
1322+
let mut observer = crate::composition::NeutralActivationObserver::new(observer);
13711323
match expert_cache.as_ref() {
13721324
Some(expert_cache) => {
13731325
let args = self.args.clone();
@@ -1410,7 +1362,7 @@ impl NemotronHModel {
14101362
cache: &mut MlxHybridState,
14111363
provider: &mut P,
14121364
stream: &Stream,
1413-
observer: &mut NeutralNemotronHObserver<'_>,
1365+
observer: &mut crate::composition::NeutralActivationObserver<'_>,
14141366
) -> Result<Array, Error>
14151367
where
14161368
P: eredu_runtime::RoutedExpertProvider<MlxNeuralBackend>,

0 commit comments

Comments
 (0)