#![forbid(unsafe_code)]
use el_core::{DeviceTarget, SafetyMode, Token};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct LogitAdjustment {
penalties: Vec<(Token, i32)>,
}
impl LogitAdjustment {
pub fn none() -> Self {
Self::default()
}
pub fn with_penalties(penalties: Vec<(Token, i32)>) -> Self {
Self { penalties }
}
pub fn is_empty(&self) -> bool {
self.penalties.is_empty()
}
pub fn delta_for(&self, token: Token) -> i32 {
self.penalties
.iter()
.find(|(t, _)| *t == token)
.map(|(_, d)| *d)
.unwrap_or(0)
}
pub fn l1_norm_milli(&self) -> u32 {
self.penalties.iter().map(|(_, d)| d.unsigned_abs()).sum()
}
}
pub trait SafetySteerer {
fn adjust(&self, recent_tokens: &[Token]) -> LogitAdjustment;
fn mode(&self) -> SafetyMode;
}
pub struct SafetyModeSelector;
impl SafetyModeSelector {
pub fn resolve(requested: SafetyMode, device: DeviceTarget) -> SafetyMode {
match (requested, device) {
(SafetyMode::SecDecoding, DeviceTarget::MidRange) => SafetyMode::Lightweight,
(m, _) => m,
}
}
}
pub struct NoSafety;
impl SafetySteerer for NoSafety {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::none()
}
fn mode(&self) -> SafetyMode {
SafetyMode::Off
}
}
pub struct LightweightFilter {
banned: Vec<Token>,
}
impl LightweightFilter {
pub const HARD_BAN: i32 = -1_000_000;
pub fn new(banned: Vec<Token>) -> Self {
Self { banned }
}
}
impl SafetySteerer for LightweightFilter {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::with_penalties(self.banned.iter().map(|&t| (t, Self::HARD_BAN)).collect())
}
fn mode(&self) -> SafetyMode {
SafetyMode::Lightweight
}
}
pub struct SecDecodingSteerer {
_private: (),
}
impl SecDecodingSteerer {
pub fn placeholder() -> Self {
Self { _private: () }
}
}
impl SafetySteerer for SecDecodingSteerer {
fn adjust(&self, _recent: &[Token]) -> LogitAdjustment {
LogitAdjustment::none()
}
fn mode(&self) -> SafetyMode {
SafetyMode::SecDecoding
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn secdecoding_downgrades_on_midrange() {
assert_eq!(
SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::MidRange),
SafetyMode::Lightweight
);
assert_eq!(
SafetyModeSelector::resolve(SafetyMode::SecDecoding, DeviceTarget::HighEnd),
SafetyMode::SecDecoding
);
}
#[test]
fn lightweight_bans_tokens() {
let f = LightweightFilter::new(vec![42, 99]);
let adj = f.adjust(&[]);
assert_eq!(adj.delta_for(42), LightweightFilter::HARD_BAN);
assert_eq!(adj.delta_for(7), 0);
assert!(adj.l1_norm_milli() > 0);
}
}