huggingface / huggingface/candle

Model Context Not Resetting Between Requests Leading to Invalid Responses

Open
#2,276 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Rust
Stars
21k
Forks
1.8k
Avg merge
16h 42m
Merged PRs (30d)
25

Description

I'm experiencing an issue with the T5 model server - which I'm trying to write as an example - where subsequent requests provide invalid responses after the first request. The first query translates correctly, but any query after that returns an incorrect response, such as `{"translated_text":"","tokens":1,"duration":0.804536125,"speed":1.2429522664380048}.`

Steps to Reproduce:
1. Setup and Run Server: Start the Rocket server with the provided code.
```
cargo run --features metal --example t5-server --release -- \
--model-id "jbochi/madlad400-7b-mt"
```
2. Initial Translation Request: Send a translation request, which returns the correct translated text.
```
curl -X POST http://localhost:8000/translate \
-H "Content-Type: application/json" \
-d '{
"text": "Hello, world!",
"translate_to": "fr"
}'
```
```
{"translated_text":" Bonjour tout le monde !","tokens":7,"duration":6.628745209,"speed":1.05600679756041}
```
3. Subsequent Requests: Send another translation request which returns:
```
{"translated_text":"","tokens":1,"duration":0.910155417,"speed":1.0987134519246617}%
```

Code: `candle-examples/examples/t5-server/main.rs`
```
use std::sync::{Arc, Mutex};
use std::path::PathBuf;
use anyhow::Result;
use candle::{DType, Device, Tensor};
use candle_transformers::models::t5;
use candle_transformers::generation::LogitsProcessor;
use candle_nn::VarBuilder;
use clap::{Parser, ValueEnum};
use hf_hub::{api::sync::Api, Repo, RepoType};
use tokenizers::Tokenizer;
use rocket::serde::{json::Json, Deserialize, Serialize};
use rocket::{post, routes, State};
use rocket::response::status::Custom;
use rocket::http::Status;

const DTYPE: DType = DType::F32;

#[derive(Clone, Debug, Copy, ValueEnum)]
enum Which {
T5Base,
T5Small,
T5Large,
T5_3B,
Mt5Base,
Mt5Small,
Mt5Large,
}

#[derive(Parser, Debug, Clone)]
#[command(author, version, about, long_about = None)]
struct Args {
#[arg(long)]
cpu: bool,

#[arg(long)]
tracing: bool,

#[arg(long)]
model_id: Option,

#[arg(long)]
revision: Option,

#[arg(long)]
model_file: Option,

#[arg(long)]
tokenizer_file: Option,

#[arg(long)]
config_file: Option,

#[arg(long)]
decode: bool,

#[arg(long, default_value = "false")]
disable_cache: bool,

#[arg(long)]
prompt: Option,

#[arg(long)]
decoder_prompt: Option,

#[arg(long, default_value = "true")]
normalize_embeddings: bool,

#[arg(long, default_value_t = 0.8)]
temperature: f64,

#[arg(long)]
top_p: Option,

#[arg(long, default_value_t = 1.1)]
repeat_penalty: f32,

#[arg(long, default_value_t = 64)]
repeat_last_n: usize,

#[arg(long, default_value = "t5-small")]
which: Which,
}
struct T5ModelBuilder {
device: Arc>,
config: t5::Config,
weights_filename: Vec,
}

impl T5ModelBuilder {
pub fn load(args: &Args) -> Result<(Self, Tokenizer)> {
let device = Arc::new(Mutex::new(candle_examples::device(args.cpu)?));
let (default_model, default_revision) = match args.which {
Which::T5Base => ("t5-base", "main"),
Which::T5Small => ("t5-small", "refs/pr/15"),
Which::T5Large => ("t5-large", "main"),
Which::T5_3B => ("t5-3b", "main"),
Which::Mt5Base => ("google/mt5-base", "refs/pr/5"),
Which::Mt5Small => ("google/mt5-small", "refs/pr/6"),
Which::Mt5Large => ("google/mt5-large", "refs/pr/2"),
};
let default_model = default_model.to_string();
let default_revision = default_revision.to_string();
let (model_id, revision) = match (args.model_id.to_owned(), args.revision.to_owned()) {
(Some(model_id), Some(revision)) => (model_id, revision),
(Some(model_id), None) => (model_id, "main".to_string()),
(None, Some(revision)) => (default_model, revision),
(None, None) => (default_model, default_revision),
};

let repo = Repo::with_revision(model_id.clone(), RepoType::Model, revision);
let api = Api::new()?;
let repo = api.repo(repo);
let config_filename = match &args.config_file {
None => repo.get("config.json")?,
Some(f) => f.into(),
};
let tokenizer_filename = match &args.tokenizer_file {
None => match args.which {
Which::Mt5Base => api
.model("lmz/mt5-tokenizers".into())
.get("mt5-base.tokenizer.json")?,
Which::Mt5Small => api
.model("lmz/mt5-tokenizers".into())
.get("mt5-small.tokenizer.json")?,
Which::Mt5Large => api
.model("lmz/mt5-tokenizers".into())
.get("mt5-large.tokenizer.json")?,
_ => repo.get("tokenizer.json")?,
},
Some(f) => f.into(),
};
let weights_filename = match &args.model_file {
Some(f) => f.split(',').map(|v| v.into()).collect::>(),
None => {
if model_id == "google/flan-t5-xxl" || model_id == "google/flan-ul2" || model_id == "jbochi/madlad400-7b-mt"{
candle_examples::hub_load_safetensors(&repo, "model.safetensors.index.json")?
} else {
vec![repo.get("model.safetensors")?]
}
}
};
let config = std::fs::read_to_string(config_filename)?;
let mut config: t5::Config = serde_json::from_str(&config)?;
config.use_cache = !args.disable_cache;
let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(|e| anyhow::anyhow!(e))?;
Ok((
Self {
device,
config,
weights_filename,
},
tokenizer,
))
}

pub fn build_conditional_generation(&self) -> Result {
let vb = unsafe {
VarBuilder::from_mmaped_safetensors(&self.weights_filename, DTYPE, &self.device.lock().unwrap())?
};
Ok(t5::T5ForConditionalGeneration::load(vb, &self.config)?)
}
}

#[derive(Deserialize)]
struct TranslationRequest {
text: String,
translate_to: String,
}

#[derive(Serialize)]
struct TranslationResponse {
translated_text: String,
tokens: usize,
duration: f64, // Duration in seconds
speed: f64, // Tokens per second
}

#[post("/translate", format = "json", data = "")]
async fn translate(
model_builder: &State,
tokenizer: &State,
model: &State>>,
device: &State>>,
request: Json
) -> Result, Custom>> {

let prompt = format!("<2{}> {}", request.translate_to, request.text);

let tokens = tokenizer
.encode(prompt, true)
.map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(e))))?
.get_ids()
.to_vec();

let input_token_ids = {
let device = device.lock().map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(format!("Failed to lock device: {}", e)))))?;
Tensor::new(&tokens[..], &*device)
.map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(e))))?
.unsqueeze(0)
.map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(e))))?
};

let mut output_token_ids = vec![model_builder.config.decoder_start_token_id.unwrap_or(model_builder.config.pad_token_id) as u32];

let encoder_output = {
let mut model = model.lock().map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(format!("Failed to lock model: {}", e)))))?;
model.encode(&input_token_ids).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?
};

let mut logits_processor = LogitsProcessor::new(299792458, Some(0.0), None);

let start = std::time::Instant::now();

for index in 0.. {
if output_token_ids.len() > 512 {
break;
}

let decoder_token_ids = if index == 0 || !model_builder.config.use_cache {
Tensor::new(output_token_ids.as_slice(), &*device.lock().unwrap()).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?.unsqueeze(0).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?
} else {
let last_token = *output_token_ids.last().unwrap();
Tensor::new(&[last_token], &*device.lock().unwrap()).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?.unsqueeze(0).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?
};

let logits = {
let mut model = model.lock().map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(anyhow::anyhow!(format!("Failed to lock model: {}", e)))))?;
model.decode(&decoder_token_ids, &encoder_output).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?.squeeze(0).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?
};

let next_token_id = logits_processor.sample(&logits).map_err(|e| Custom(Status::InternalServerError, rocket::response::Debug(e.into())))?;
if next_token_id as usize == model_builder.config.eos_token_id {
break;
}
output_token_ids.push(next_token_id);
}

let dt = start.elapsed().as_secs_f64();

let translated_text: String = output_token_ids.iter()
.map(|&tok| tokenizer.id_to_token(tok).unwrap_or_else(|| "".to_string()).replace('▁', " ").replace("<0x0A>", "\n"))
.collect::>()
.join("");

Ok(Json(TranslationResponse {
translated_text,
tokens: output_token_ids.len(),
duration: dt,
speed: output_token_ids.len() as f64 / dt,
}))
}

#[rocket::main]
async fn main() -> Result<()> {
use tracing_chrome::ChromeLayerBuilder;
use tracing_subscriber::prelude::*;

let args = Args::parse();

let _guard = if args.tracing {
let (chrome_layer, guard) = ChromeLayerBuilder::new().build();
tracing_subscriber::registry().with(chrome_layer).init();
Some(guard)
} else {
None
};

let (builder, tokenizer) = T5ModelBuilder::load(&args)?;
let model = Arc::new(Mutex::new(builder.build_conditional_generation()?));
let device = builder.device.clone();

let _rocket = rocket::build()
.mount("/", routes![translate])
.manage(builder)
.manage(tokenizer)
.manage(model)
.manage(device)
.launch()
.await?;

Ok(())
}
pub fn normalize_l2(v: &Tensor) -> Result {
Ok(v.broadcast_div(&v.sqr()?.sum_keepdim(1)?.sqrt()?)?)
}

```

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.