use std::{
collections::HashMap,
path::{Path, PathBuf},
sync::Arc,
};
use anyhow::{bail, Result};
use derivative::Derivative;
use flume::{Receiver, Sender};
use futures::future::join_all;
use half::f16;
use itertools::Itertools;
use memmap2::Mmap;
use reload::{AdapterOption, BnfOption, Precision};
use safetensors::SafeTensors;
use salvo::oapi::ToSchema;
use serde::{de::DeserializeSeed, Deserialize, Serialize};
use tokio::{
fs::File,
io::{AsyncReadExt, BufReader},
sync::RwLock,
time::Duration,
};
use web_rwkv::{
context::{Context, ContextBuilder, ContextError, InstanceExt},
runtime::{
infer::Rnn,
loader::{Loader, Lora, LoraBlend, Reader},
model::{Bundle, ContextAutoLimits, ModelBuilder, ModelInfo, ModelVersion, Quant, State},
v4, v5, v6, v7, Runtime, TokioRuntime,
},
tensor::{serialization::Seed, TensorCpu, TensorError, TensorInit},
tokenizer::Tokenizer,
wgpu::{Backends, PowerPreference},
};
use crate::{run::GenerateContext, sampler::Sampler};
pub mod reload;
pub mod run;
pub mod sampler;
pub const MAX_TOKENS: usize = usize::MAX;
#[derive(Debug)]
pub enum Token {
Start,
Content(String),
Stop(FinishReason, TokenCounter),
Embed(Vec<f32>, [usize; 4]),
Choose(Vec<f32>),
Done,
}
#[derive(Debug, Default, Clone, Serialize, Deserialize, ToSchema)]
pub struct TokenCounter {
#[serde(alias = "prompt_tokens")]
pub prompt: usize,
#[serde(alias = "completion_tokens")]
pub completion: usize,
#[serde(alias = "total_tokens")]
pub total: usize,
pub duration: Duration,
}
#[derive(Debug, Default, Clone, Copy, Serialize, ToSchema)]
#[serde(rename_all = "snake_case")]
#[allow(dead_code)]
pub enum FinishReason {
Stop,
Length,
ContentFilter,
#[default]
#[serde(untagged)]
Null,
}
#[derive(Debug, Clone)]
pub enum ThreadRequest {
Adapter(Sender<AdapterList>),
Info(Sender<RuntimeInfo>),
Generate {
request: Box<GenerateRequest>,
tokenizer: Arc<Tokenizer>,
sender: Sender<Token>,
},
Reload {
request: Box<ReloadRequest>,
sender: Option<Sender<bool>>,
},
Unload,
Save {
request: SaveRequest,
sender: Sender<bool>,
},
}
#[derive(Default)]
pub enum Environment {
Loaded {
info: RuntimeInfo,
runtime: Arc<dyn Runtime<Rnn> + Send + Sync>,
model: Arc<dyn ModelSerialize + Send + Sync>,
sender: Sender<GenerateContext>,
},
#[default]
None,
}
#[derive(Derivative, Clone)]
#[derivative(Debug)]
pub struct RuntimeInfo {
pub reload: Arc<ReloadRequest>,
pub info: ModelInfo,
pub states: Vec<InitState>,
pub tokenizer: Arc<Tokenizer>,
}
struct Model<M>(M);
pub trait ModelSerialize {
fn serialize(&self, file: std::fs::File) -> Result<()>;
}
impl<M: Serialize> ModelSerialize for Model<M> {
fn serialize(&self, file: std::fs::File) -> Result<()> {
use cbor4ii::{core::enc::Write, serde::Serializer};
use std::{fs::File, io::Write as _};
struct FileWriter(File);
impl Write for FileWriter {
type Error = std::io::Error;
fn push(&mut self, input: &[u8]) -> Result<(), Self::Error> {
self.0.write_all(input)
}
}
let file = FileWriter(file);
let mut serializer = Serializer::new(file);
self.0.serialize(&mut serializer)?;
Ok(())
}
}
#[derive(Debug, Default, Clone)]
pub struct AdapterList(pub Vec<String>);
#[derive(Debug, Default, Clone)]
pub enum GenerateKind {
#[default]
None,
State,
Choose {
choices: Vec<String>,
calibrate: bool,
},
}
#[derive(Clone, Derivative)]
#[derivative(Debug, Default)]
pub struct GenerateRequest {
pub prompt: String,
pub model_text: String,
pub max_tokens: usize,
pub stop: Vec<String>,
pub bias: Arc<HashMap<u32, f32>>,
pub bnf_schema: Option<String>,
#[derivative(
Debug = "ignore",
Default(value = "Arc::new(RwLock::new(sampler::nucleus::NucleusSampler::default()))")
)]
pub sampler: Arc<RwLock<dyn Sampler + Send + Sync>>,
pub kind: GenerateKind,
pub state: Arc<InputState>,
}
#[derive(Debug, Derivative, Clone, Serialize, Deserialize, ToSchema)]
#[derivative(Default)]
#[serde(default)]
pub struct ReloadRequest {
#[salvo(schema(value_type = String))]
pub model_path: PathBuf,
pub lora: Vec<reload::Lora>,
pub state: Vec<reload::State>,
pub quant: usize,
#[salvo(schema(value_type = sealed::Quant))]
pub quant_type: Quant,
pub precision: Precision,
#[derivative(Default(value = "128"))]
pub token_chunk_size: usize,
#[derivative(Default(value = "8"))]
pub max_batch: usize,
#[salvo(schema(value_type = String))]
pub tokenizer_path: PathBuf,
pub bnf: BnfOption,
pub adapter: AdapterOption,
}
#[derive(Debug, Default, Clone, Serialize, Deserialize, ToSchema)]
#[serde(default)]
pub struct SaveRequest {
#[serde(alias = "model_path")]
#[salvo(schema(value_type = String))]
pub path: PathBuf,
}
#[derive(Debug, Deserialize)]
struct Prefab {
info: ModelInfo,
}
#[derive(Debug, Clone, Copy)]
enum LoadType {
SafeTensors,
Prefab,
}
#[derive(
Derivative, Default, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToSchema,
)]
#[derivative(Debug = "transparent")]
#[serde(transparent)]
pub struct StateId(uuid::Uuid);
impl StateId {
pub fn new() -> Self {
Self(uuid::Uuid::new_v4())
}
}
#[derive(Debug, Default, Clone, Serialize, Deserialize, ToSchema)]
pub struct StateValue {
pub name: String,
pub id: StateId,
pub data: Vec<f32>,
pub shape: [usize; 4],
}
#[derive(Debug, Default, Clone, Serialize, Deserialize, ToSchema)]
pub struct StateFile {
pub name: String,
pub id: StateId,
#[salvo(schema(value_type = String))]
pub path: PathBuf,
}
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)]
#[serde(untagged)]
pub enum InputState {
Key(StateId),
Value(StateValue),
File(StateFile),
}
impl Default for InputState {
fn default() -> Self {
Self::Key(Default::default())
}
}
impl InputState {
pub fn id(&self) -> StateId {
match self {
InputState::Key(id) => *id,
InputState::Value(value) => value.id,
InputState::File(file) => file.id,
}
}
}
#[derive(Derivative, Clone, Serialize, Deserialize)]
#[derivative(Debug)]
pub struct InitState {
pub name: String,
pub id: StateId,
pub default: bool,
#[derivative(Debug = "ignore")]
pub data: TensorCpu<f32>,
}
impl TryFrom<StateValue> for InitState {
type Error = TensorError;
fn try_from(
StateValue {
name,
id,
data,
shape,
}: StateValue,
) -> Result<Self, Self::Error> {
let default = false;
let data = TensorCpu::from_data(shape, data)?;
Ok(Self {
name,
id,
default,
data,
})
}
}
fn list_adapters() -> AdapterList {
let backends = Backends::all();
let instance = web_rwkv::wgpu::Instance::default();
let list = instance
.enumerate_adapters(backends)
.into_iter()
.map(|adapter| adapter.get_info())
.map(|info| format!("{} ({:?})", info.name, info.backend))
.collect();
AdapterList(list)
}
async fn create_context(adapter: AdapterOption, info: &ModelInfo) -> Result<Context> {
let backends = Backends::all();
let instance = web_rwkv::wgpu::Instance::default();
let adapter = match adapter {
AdapterOption::Auto => instance.adapter(PowerPreference::HighPerformance).await,
AdapterOption::Economical => instance.adapter(PowerPreference::LowPower).await,
AdapterOption::Manual(selection) => Ok(instance
.enumerate_adapters(backends)
.into_iter()
.nth(selection)
.ok_or(ContextError::RequestAdapterFailed)?),
}?;
let context = ContextBuilder::new(adapter)
.auto_limits(info)
.build()
.await?;
Ok(context)
}
async fn load_tokenizer(path: impl AsRef<Path>) -> Result<Tokenizer> {
let file = File::open(path).await?;
let mut reader = BufReader::new(file);
let mut contents = String::new();
reader.read_to_string(&mut contents).await?;
Ok(Tokenizer::new(&contents)?)
}
async fn load_model_state<R: Reader>(
context: &Context,
info: &ModelInfo,
model: R,
) -> Result<TensorCpu<f32>> {
match info.version {
ModelVersion::V4 => bail!("v4 does not support init state yet"),
ModelVersion::V5 => Ok(v5::read_state(context, info, model).await?),
ModelVersion::V6 => Ok(v6::read_state(context, info, model).await?),
ModelVersion::V7 => Ok(v7::read_state(context, info, model).await?),
}
}
async fn load_runtime(
context: &Context,
info: &ModelInfo,
request: &ReloadRequest,
load: LoadType,
) -> Result<(
Vec<InitState>,
Arc<dyn Runtime<Rnn> + Send + Sync>,
Arc<dyn State + Send + Sync>,
Arc<dyn ModelSerialize + Send + Sync>,
)> {
let ReloadRequest {
model_path,
lora,
state,
quant,
quant_type,
precision,
max_batch,
..
} = request.clone();
let mut states = Vec::with_capacity(state.len());
for state in state.into_iter() {
let reload::State {
path,
name,
id,
default,
} = state;
let name = match name {
Some(name) => name,
None => match path.file_name() {
Some(name) => name.to_string_lossy().to_string(),
None => continue,
},
};
let file = File::open(path).await?;
let data = unsafe { Mmap::map(&file) }?;
let model = SafeTensors::deserialize(&data)?;
match load_model_state(context, info, model).await {
Ok(data) => {
let state = InitState {
name,
id,
data,
default,
};
log::info!("{:#?}", state);
states.push(state);
}
Err(err) => log::warn!("initial state not loaded: {}", err),
}
}
let file = File::open(model_path).await?;
let data = unsafe { Mmap::map(&file) }?;
match load {
LoadType::SafeTensors => {
let model = SafeTensors::deserialize(&data)?;
if let Ok(data) = load_model_state(context, info, model).await {
let name = "internal".into();
let id = StateId::new();
let state = InitState {
name,
id,
data,
default: true,
};
states.push(state);
}
let model = SafeTensors::deserialize(&data)?;
let quant = (0..quant).map(|layer| (layer, quant_type)).collect();
let lora: Vec<Result<_>> = join_all(lora.iter().map(|lora| async move {
let reload::Lora { path, alpha } = lora;
let file = File::open(path).await?;
let data = unsafe { Mmap::map(&file)? };
let blend = LoraBlend::full(*alpha);
Ok((data, blend))
}))
.await;
let lora: Vec<_> = lora.into_iter().try_collect()?;
let lora: Vec<_> = lora
.iter()
.map(|(data, blend)| -> Result<_> {
let data = SafeTensors::deserialize(data)?;
let blend = blend.clone();
Ok(Lora { data, blend })
})
.try_collect()?;
let builder = ModelBuilder::new(context, model).quant(quant);
let builder = lora.into_iter().fold(builder, |builder, x| builder.lora(x));
macro_rules! match_safe_tensors {
(($v:expr, $p:expr), { $(($version:path, $precision:path, $model:ty, $build:ident, $bundle:ty)),+ }) => {
match ($v, $p) {
$(
($version, $precision) => {
let model = builder.$build().await?;
let bundle = <$bundle>::new(model, max_batch);
let state = Arc::new(bundle.state());
let model = Arc::new(Model(bundle.model()));
let runtime = Arc::new(TokioRuntime::<Rnn>::new(bundle).await);
Ok((states, runtime, state, model))
}
)+
}
}
}
match_safe_tensors!(
(info.version, precision),
{
(ModelVersion::V4, Precision::Fp16, v4::Model, build_v4, v4::Bundle::<f16>),
(ModelVersion::V5, Precision::Fp16, v5::Model, build_v5, v5::Bundle::<f16>),
(ModelVersion::V6, Precision::Fp16, v6::Model, build_v6, v6::Bundle::<f16>),
(ModelVersion::V7, Precision::Fp16, v7::Model, build_v7, v7::Bundle::<f16>),
(ModelVersion::V4, Precision::Fp32, v4::Model, build_v4, v4::Bundle::<f32>),
(ModelVersion::V5, Precision::Fp32, v5::Model, build_v5, v5::Bundle::<f32>),
(ModelVersion::V6, Precision::Fp32, v6::Model, build_v6, v6::Bundle::<f32>),
(ModelVersion::V7, Precision::Fp32, v7::Model, build_v7, v7::Bundle::<f32>)
}
)
}
LoadType::Prefab => {
use cbor4ii::{core::utils::SliceReader, serde::Deserializer};
let reader = SliceReader::new(&data);
let mut deserializer = Deserializer::new(reader);
macro_rules! match_prefab {
(($v:expr, $p:expr), { $(($version:path, $precision:path, $model:ty, $bundle:ty)),+ }) => {
match ($v, $p) {
$(
($version, $precision) => {
let seed: Seed<_, $model> = Seed::new(context);
let model = seed.deserialize(&mut deserializer)?;
let bundle = <$bundle>::new(model, max_batch);
let state = Arc::new(bundle.state());
let model = Arc::new(Model(bundle.model()));
let runtime = Arc::new(TokioRuntime::<Rnn>::new(bundle).await);
Ok((states, runtime, state, model))
}
)+
}
}
}
match_prefab!(
(info.version, precision),
{
(ModelVersion::V4, Precision::Fp16, v4::Model, v4::Bundle::<f16>),
(ModelVersion::V5, Precision::Fp16, v5::Model, v5::Bundle::<f16>),
(ModelVersion::V6, Precision::Fp16, v6::Model, v6::Bundle::<f16>),
(ModelVersion::V7, Precision::Fp16, v7::Model, v7::Bundle::<f16>),
(ModelVersion::V4, Precision::Fp32, v4::Model, v4::Bundle::<f32>),
(ModelVersion::V5, Precision::Fp32, v5::Model, v5::Bundle::<f32>),
(ModelVersion::V6, Precision::Fp32, v6::Model, v6::Bundle::<f32>),
(ModelVersion::V7, Precision::Fp32, v7::Model, v7::Bundle::<f32>)
}
)
}
}
}
async fn process(env: Arc<RwLock<Environment>>, request: ThreadRequest) -> Result<()> {
match request {
ThreadRequest::Adapter(sender) => {
let _ = sender.send(list_adapters());
}
ThreadRequest::Info(sender) => {
let env = env.read().await;
if let Environment::Loaded { info, .. } = &*env {
let _ = sender.send(info.clone());
}
}
ThreadRequest::Generate {
request,
tokenizer,
sender,
} => {
let context = GenerateContext::new(*request, sender, &tokenizer).await?;
let env = env.read().await;
if let Environment::Loaded { sender, .. } = &*env {
let _ = sender.send(context);
}
}
ThreadRequest::Reload { request, sender } => {
let handle = tokio::spawn(async move {
let file = File::open(&request.model_path).await?;
let data = unsafe { Mmap::map(&file)? };
let (info, load) = {
let st = SafeTensors::deserialize(&data);
let prefab = cbor4ii::serde::from_slice::<Prefab>(&data);
match (st, prefab) {
(Ok(model), _) => (Loader::info(&model)?, LoadType::SafeTensors),
(_, Ok(prefab)) => (prefab.info, LoadType::Prefab),
_ => bail!("failed to read model info"),
}
};
log::info!("{:#?}", request);
log::info!("{:#?}", info);
log::info!("model type: {:?}", load);
let context = create_context(request.adapter, &info).await?;
log::info!("{:#?}", context.adapter.get_info());
let mut env = env.write().await;
let _ = std::mem::take(&mut *env);
let tokenizer = Arc::new(load_tokenizer(&request.tokenizer_path).await?);
let (states, runtime, state, model) =
load_runtime(&context, &info, &request, load).await?;
let reload = Arc::new(*request);
let info = RuntimeInfo {
reload,
info,
states,
tokenizer,
};
let sender = {
let runtime = Arc::downgrade(&runtime);
let (sender, receiver) = flume::unbounded();
tokio::spawn(crate::run::run(
context,
runtime,
state,
receiver,
info.clone(),
));
sender
};
log::info!("model loaded");
let _ = std::mem::replace(
&mut *env,
Environment::Loaded {
info,
runtime,
model,
sender,
},
);
Ok(())
});
if let Some(sender) = sender {
let _ = match handle.await? {
Ok(_) => sender.send(true),
Err(err) => {
log::error!("[reload] error: {err:#?}");
sender.send(false)
}
};
}
}
ThreadRequest::Unload => {
let mut env = env.write().await;
let _ = std::mem::take(&mut *env);
log::info!("model unloaded");
}
ThreadRequest::Save { request, sender } => {
let env = env.read().await;
if let Environment::Loaded { model, .. } = &*env {
log::info!("serializing model into {:?}", &request.path);
let model = model.clone();
let handle = tokio::task::spawn_blocking(move || {
let file = std::fs::File::create(request.path)?;
model.serialize(file)
});
drop(env);
let _ = match handle.await? {
Ok(_) => sender.send(true),
Err(err) => {
log::error!("[save] error: {err:#?}");
sender.send(false)
}
};
}
}
};
Ok(())
}
pub async fn serve(receiver: Receiver<ThreadRequest>) {
let env: Arc<RwLock<Environment>> = Default::default();
while let Ok(request) = receiver.recv_async().await {
let future = process(env.clone(), request);
tokio::spawn(future);
}
}
#[allow(dead_code)]
mod sealed {
use salvo::oapi::ToSchema;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, ToSchema)]
pub enum Quant {
#[default]
None,
Int8,
NF4,
SF4,
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, ToSchema)]
pub enum EmbedDevice {
#[default]
Cpu,
Gpu,
}
}