1#![forbid(unsafe_code)]
10
11use el_core::{DeviceTarget, SafetyMode, Token};
12
13#[derive(Debug, Clone, Default, PartialEq, Eq)]
16pub struct LogitAdjustment {
17 penalties: Vec<(Token, i32)>,
18}
19
20impl LogitAdjustment {
21 pub fn none() -> Self {
22 Self::default()
23 }
24
25 pub fn with_penalties(penalties: Vec<(Token, i32)>) -> Self {
26 Self { penalties }
27 }
28
29 pub fn is_empty(&self) -> bool {
30 self.penalties.is_empty()
31 }
32
33 pub fn delta_for(&self, token: Token) -> i32 {
35 self.penalties
36 .iter()
37 .find(|(t, _)| *t == token)
38 .map(|(_, d)| *d)
39 .unwrap_or(0)
40 }
41
42 pub fn l1_norm_milli(&self) -> u32 {
45 self.penalties.iter().map(|(_, d)| d.unsigned_abs()).sum()
46 }
47}
48
49pub trait SafetySteerer {
52 fn adjust(&self, recent_tokens: &[Token]) -> LogitAdjustment;
53 fn mode(&self) -> SafetyMode;
54}
55
56pub struct SafetyModeSelector;
58
59impl SafetyModeSelector {
60 pub fn resolve(requested: SafetyMode, device: DeviceTarget) -> SafetyMode {
63 match (requested, device) {
64 (SafetyMode::SecDecoding, DeviceTarget::MidRange) => SafetyMode::Lightweight,
65 (m, _) => m,
66 }
67 }
68}
69
70pub struct NoSafety;
72
73impl SafetySteerer for NoSafety {
74 fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
75 LogitAdjustment::none()
76 }
77 fn mode(&self) -> SafetyMode {
78 SafetyMode::Off
79 }
80}
81
82pub struct LightweightFilter {
85 banned: Vec<Token>,
86}
87
88impl LightweightFilter {
89 pub const HARD_BAN: i32 = -1_000_000;
90
91 pub fn new(banned: Vec<Token>) -> Self {
92 Self { banned }
93 }
94}
95
96impl SafetySteerer for LightweightFilter {
97 fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
98 LogitAdjustment::with_penalties(self.banned.iter().map(|&t| (t, Self::HARD_BAN)).collect())
99 }
100 fn mode(&self) -> SafetyMode {
101 SafetyMode::Lightweight
102 }
103}
104
105pub struct SecDecodingSteerer {
111 _private: (),
112}
113
114impl SecDecodingSteerer {
115 pub fn placeholder() -> Self {
116 Self { _private: () }
117 }
118}
119
120impl SafetySteerer for SecDecodingSteerer {
121 fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
122 LogitAdjustment::none()
125 }
126 fn mode(&self) -> SafetyMode {
127 SafetyMode::SecDecoding
128 }
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134
135 #[test]
136 fn secdecoding_downgrades_on_midrange() {
137 assert_eq!(
138 SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::MidRange),
139 SafetyMode::Lightweight
140 );
141 assert_eq!(
142 SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::HighEnd),
143 SafetyMode::SecDecoding
144 );
145 }
146
147 #[test]
148 fn lightweight_bans_tokens() {
149 let f = LightweightFilter::new(vec![42, 99]);
150 let adj = f.adjust(&[]);
151 assert_eq!(adj.delta_for(42), LightweightFilter::HARD_BAN);
152 assert_eq!(adj.delta_for(7), 0);
153 assert!(adj.l1_norm_milli() > 0);
154 }
155}