use std::collections::HashMap;
use dynamo_runtime::protocols::annotated::AnnotationsProvider;
use serde::{Deserialize, Serialize};
use utoipa::ToSchema;
use validator::Validate;
mod aggregator;
use super::PromptTruncationSide;
pub use super::embeddings::{NvExt, NvExtProvider};
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
#[serde(untagged)]
pub enum PoolingInput {
Single(String),
Batch(Vec<String>),
Tokens(Vec<u32>),
TokenBatch(Vec<Vec<u32>>),
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq, Default)]
#[serde(rename_all = "snake_case")]
pub enum PoolingEncodingFormat {
#[default]
Float,
Base64,
Bytes,
BytesOnly,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum PoolingEmbedDType {
#[default]
#[serde(rename = "float32")]
Float32,
#[serde(rename = "float16")]
Float16,
#[serde(rename = "bfloat16")]
Bfloat16,
#[serde(rename = "fp8_e4m3")]
Fp8E4m3,
#[serde(rename = "fp8_e5m2")]
Fp8E5m2,
}
impl PoolingEmbedDType {
pub fn as_str(self) -> &'static str {
match self {
Self::Float32 => "float32",
Self::Float16 => "float16",
Self::Bfloat16 => "bfloat16",
Self::Fp8E4m3 => "fp8_e4m3",
Self::Fp8E5m2 => "fp8_e5m2",
}
}
pub(crate) const fn byte_width(self) -> usize {
match self {
Self::Float32 => 4,
Self::Float16 | Self::Bfloat16 => 2,
Self::Fp8E4m3 | Self::Fp8E5m2 => 1,
}
}
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq, Default)]
#[serde(rename_all = "snake_case")]
pub enum PoolingEndianness {
#[default]
Native,
Big,
Little,
}
impl PoolingEndianness {
pub fn as_str(self) -> &'static str {
match self {
Self::Native => "native",
Self::Big => "big",
Self::Little => "little",
}
}
}
#[derive(ToSchema, Serialize, Deserialize, Validate, Debug, Clone)]
pub struct NvCreatePoolingRequest {
pub model: String,
pub input: PoolingInput,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
#[serde(default)]
pub priority: i64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mm_processor_kwargs: Option<HashMap<String, serde_json::Value>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_salt: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub task: Option<String>,
#[serde(default)]
pub encoding_format: PoolingEncodingFormat,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub use_activation: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub add_special_tokens: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub truncation_side: Option<PromptTruncationSide>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dimensions: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub embed_dtype: Option<PoolingEmbedDType>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endianness: Option<PoolingEndianness>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub truncate_prompt_tokens: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub nvext: Option<NvExt>,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
#[serde(untagged)]
pub enum PoolingOutput {
Base64(String),
Vector(Vec<f32>),
Matrix(Vec<Vec<f32>>),
}
fn default_pooling_object() -> String {
"pooling".to_string()
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone)]
pub struct PoolingData {
pub index: u32,
#[serde(default = "default_pooling_object")]
pub object: String,
pub data: PoolingOutput,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[schema(ignore)]
pub shape: Option<Vec<u64>>,
}
#[derive(ToSchema, Serialize, Deserialize, Debug, Clone, Default)]
pub struct PoolingUsage {
pub prompt_tokens: u32,
pub total_tokens: u32,
#[serde(default)]
pub completion_tokens: u32,
}
#[derive(ToSchema, Serialize, Deserialize, Validate, Debug, Clone)]
pub struct NvCreatePoolingResponse {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub data: Vec<PoolingData>,
pub usage: PoolingUsage,
}
impl NvCreatePoolingResponse {
pub fn empty() -> Self {
Self {
id: String::new(),
object: "list".to_string(),
created: 0,
model: "pooling".to_string(),
data: vec![],
usage: PoolingUsage::default(),
}
}
}
impl NvExtProvider for NvCreatePoolingRequest {
fn nvext(&self) -> Option<&NvExt> {
self.nvext.as_ref()
}
}
impl AnnotationsProvider for NvCreatePoolingRequest {
fn annotations(&self) -> Option<Vec<String>> {
self.nvext
.as_ref()
.and_then(|nvext| nvext.annotations.clone())
}
fn has_annotation(&self, annotation: &str) -> bool {
self.nvext
.as_ref()
.and_then(|nvext| nvext.annotations.as_ref())
.map(|annotations| annotations.contains(&annotation.to_string()))
.unwrap_or(false)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn single_input_round_trips_and_wire_carries_encoding_format() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello world"
}))
.unwrap();
assert!(matches!(request.input, PoolingInput::Single(_)));
assert_eq!(request.encoding_format, PoolingEncodingFormat::Float);
assert!(request.task.is_none());
let value = serde_json::to_value(&request).unwrap();
assert_eq!(value["encoding_format"], "float");
assert!(value.get("task").is_none());
}
#[test]
fn token_inputs_parse_as_token_variants() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": [101, 2023, 102]
}))
.unwrap();
assert!(matches!(request.input, PoolingInput::Tokens(_)));
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": [[101, 102], [101, 103]]
}))
.unwrap();
match &request.input {
PoolingInput::TokenBatch(v) => assert_eq!(v.len(), 2),
other => panic!("expected token batch, got {other:?}"),
}
}
#[test]
fn task_and_options_round_trip() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": ["a", "b"],
"task": "token_classify",
"encoding_format": "base64",
"use_activation": false
}))
.unwrap();
assert_eq!(request.task.as_deref(), Some("token_classify"));
assert_eq!(request.encoding_format, PoolingEncodingFormat::Base64);
assert_eq!(request.use_activation, Some(false));
}
#[test]
fn binary_encoding_options_round_trip() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hi",
"encoding_format": "bytes",
"embed_dtype": "float16",
"endianness": "big"
}))
.unwrap();
assert_eq!(request.encoding_format, PoolingEncodingFormat::Bytes);
assert_eq!(request.embed_dtype, Some(PoolingEmbedDType::Float16));
assert_eq!(request.endianness, Some(PoolingEndianness::Big));
let value = serde_json::to_value(&request).unwrap();
assert_eq!(value["encoding_format"], "bytes");
assert_eq!(value["embed_dtype"], "float16");
assert_eq!(value["endianness"], "big");
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hi",
"encoding_format": "bytes_only",
"embed_dtype": "fp8_e4m3",
"endianness": "little"
}))
.unwrap();
assert_eq!(request.encoding_format, PoolingEncodingFormat::BytesOnly);
assert_eq!(request.embed_dtype, Some(PoolingEmbedDType::Fp8E4m3));
assert_eq!(request.endianness, Some(PoolingEndianness::Little));
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hi"
}))
.unwrap();
assert!(request.embed_dtype.is_none() && request.endianness.is_none());
let value = serde_json::to_value(&request).unwrap();
assert!(value.get("embed_dtype").is_none() && value.get("endianness").is_none());
}
#[test]
fn tokenization_options_round_trip() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hi",
"add_special_tokens": false,
"truncate_prompt_tokens": 64,
"truncation_side": "left"
}))
.unwrap();
assert_eq!(request.add_special_tokens, Some(false));
assert_eq!(request.truncate_prompt_tokens, Some(64));
assert_eq!(request.truncation_side, Some(PromptTruncationSide::Left));
}
#[test]
fn pooling_request_controls_round_trip() {
let request: NvCreatePoolingRequest = serde_json::from_value(json!({
"model": "test-model",
"input": "hello",
"user": "user-1",
"request_id": "request-1",
"priority": -2,
"mm_processor_kwargs": {"do_resize": false},
"cache_salt": "salt"
}))
.unwrap();
assert_eq!(request.user.as_deref(), Some("user-1"));
assert_eq!(request.request_id.as_deref(), Some("request-1"));
assert_eq!(request.priority, -2);
assert_eq!(request.cache_salt.as_deref(), Some("salt"));
let value = serde_json::to_value(request).unwrap();
assert_eq!(value["mm_processor_kwargs"]["do_resize"], false);
assert_eq!(value["priority"], -2);
}
#[test]
fn response_data_accepts_vector_matrix_and_base64() {
let response: NvCreatePoolingResponse = serde_json::from_value(json!({
"id": "pool-1",
"object": "list",
"created": 1,
"model": "m",
"data": [
{"index": 0, "object": "pooling", "data": [0.1, 0.2]},
{"index": 1, "object": "pooling", "data": [[0.1], [0.2]]},
{"index": 2, "object": "pooling", "data": "AAAA"}
],
"usage": {"prompt_tokens": 3, "total_tokens": 3}
}))
.unwrap();
assert!(matches!(response.data[0].data, PoolingOutput::Vector(_)));
assert!(matches!(response.data[1].data, PoolingOutput::Matrix(_)));
assert!(matches!(response.data[2].data, PoolingOutput::Base64(_)));
assert_eq!(response.usage.completion_tokens, 0);
}
}