From 351c108f3c3831259138b08a0adecb582cd436ef Mon Sep 17 00:00:00 2001 From: Alvaro Bartolome <36760800+alvarobartt@users.noreply.github.com> Date: Fri, 25 Sep 2026 13:20:50 +0000 Subject: [PATCH] Add safety note around `unsafe` for `mmap2` --- src/models/laya.rs | 29 ++++++++++++++++------------- src/models/modernbert.rs | 9 +++++++-- 2 files changed, 23 insertions(+), 15 deletions(-) diff --git a/src/models/laya.rs b/src/models/laya.rs index f214926..751d025 100644 --- a/src/models/laya.rs +++ b/src/models/laya.rs @@ -1,8 +1,8 @@ -use super::AttentionImplementation; -use super::DecisionModel; use super::modernbert::{ AttentionOptions, Config as ModernBertConfig, Encoder as ModernBertEncoder, }; +use super::AttentionImplementation; +use super::DecisionModel; use crate::{ device, schema::{ApiError, DecisionRequest, DecisionResponse, Usage}, @@ -10,10 +10,10 @@ use crate::{ }; use anyhow::Context; -use candle_core::{D, DType, Device, IndexOp, Tensor}; -use candle_nn::{Embedding, LayerNorm, Linear, VarBuilder, embedding, layer_norm}; +use candle_core::{DType, Device, IndexOp, Tensor, D}; +use candle_nn::{embedding, layer_norm, Embedding, LayerNorm, Linear, VarBuilder}; use serde::Deserialize; -use serde_json::{Map, Value, json}; +use serde_json::{json, Map, Value}; use std::{collections::HashMap, fs, path::Path}; const TYPES: [&str; 3] = ["choice", "score", "noul"]; @@ -201,6 +201,11 @@ impl Laya { config.max_len = max_model_len; } let weights = path.join("model.safetensors"); + // SAFETY: All file-backed mmap constructors are marked `unsafe` because of the potential + // for undefined behavior using the map if the underlying file is modified or out of + // process. + // + // More information at https://github.com/RazrFalcon/memmap2-rs/blob/a02e2a/src/lib.rs#L135-L165 let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights], model_dtype, &device)? }; let encoder_vb = vb.clone().rename_f(|name| { name.strip_prefix("model.") @@ -1020,14 +1025,12 @@ mod tests { #[test] fn rejects_flash_attention_without_a_compatible_backend() { - assert!( - validate_attention( - AttentionImplementation::FlashAttention2, - &Device::Cpu, - DType::BF16, - ) - .is_err() - ); + assert!(validate_attention( + AttentionImplementation::FlashAttention2, + &Device::Cpu, + DType::BF16, + ) + .is_err()); validate_attention(AttentionImplementation::Eager, &Device::Cpu, DType::F32).unwrap(); } diff --git a/src/models/modernbert.rs b/src/models/modernbert.rs index 84c21db..026a980 100644 --- a/src/models/modernbert.rs +++ b/src/models/modernbert.rs @@ -1,7 +1,7 @@ use super::AttentionImplementation; -use candle_core::{D, DType, Device, Result, Tensor}; +use candle_core::{DType, Device, Result, Tensor, D}; use candle_nn::{ - Embedding, LayerNorm, Linear, Module, VarBuilder, embedding, layer_norm_no_bias, ops::softmax, + embedding, layer_norm_no_bias, ops::softmax, Embedding, LayerNorm, Linear, Module, VarBuilder, }; use serde::Deserialize; use std::{ @@ -669,6 +669,11 @@ mod tests { let device = crate::device::load()?; let config = Config::load(&path.join("encoder/config.json"))?; let weights = path.join("model.safetensors"); + // SAFETY: All file-backed mmap constructors are marked `unsafe` because of the potential + // for undefined behavior using the map if the underlying file is modified or out of + // process. + // + // More information at https://github.com/RazrFalcon/memmap2-rs/blob/a02e2a/src/lib.rs#L135-L165 let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights], DType::F32, &device)? }; let vb = vb.rename_f(|name| { name.strip_prefix("model.")