use crate::error::Result;
use crate::models::sequential::Sequential;
use crate::models::Model;
use scirs2_core::ndarray::ArrayD;
use scirs2_core::numeric::{Float, FromPrimitive, NumAssign};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fmt::{Debug, Display};
use std::path::PathBuf;
#[derive(Debug, Clone)]
pub struct PackageResult {
pub format: PackageFormat,
pub platform: TargetPlatform,
pub output_paths: Vec<PathBuf>,
pub metadata: PackageMetadata,
}
#[derive(Debug, Clone)]
pub struct MobileConfig {
pub platform: MobilePlatform,
pub min_os_version: String,
pub architecture: MobileArchitecture,
pub optimization: MobileOptimization,
pub framework_config: FrameworkConfig,
}
#[derive(Debug, Clone)]
pub struct ServerConfig {
pub port: u16,
pub max_batch_size: usize,
pub timeout_seconds: u64,
pub enable_logging: bool,
pub max_concurrent_requests: usize,
}
#[derive(Debug, Clone)]
pub struct MobileOptimization {
pub enable_quantization: bool,
pub pruning_level: f64,
pub memory_optimization: bool,
pub battery_optimization: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum TargetPlatform {
LinuxX64,
LinuxArm64,
WindowsX64,
MacOSX64,
MacOSArm64,
AndroidArm64,
AndroidX64,
IOSArm64,
IOSX64,
WASM,
}
impl TargetPlatform {
pub fn host() -> Option<Self> {
match (std::env::consts::OS, std::env::consts::ARCH) {
("linux", "x86_64") => Some(Self::LinuxX64),
("linux", "aarch64") => Some(Self::LinuxArm64),
("windows", "x86_64") => Some(Self::WindowsX64),
("macos", "x86_64") => Some(Self::MacOSX64),
("macos", "aarch64") => Some(Self::MacOSArm64),
_ => None,
}
}
pub fn is_host(&self) -> bool {
TargetPlatform::host().as_ref() == Some(self)
}
pub(super) fn exe_suffix(&self) -> &'static str {
match self {
TargetPlatform::WindowsX64 => ".exe",
_ => "",
}
}
pub(super) fn dylib_prefix_suffix(&self) -> (&'static str, &'static str) {
match self {
TargetPlatform::WindowsX64 => ("", ".dll"),
TargetPlatform::MacOSX64 | TargetPlatform::MacOSArm64 => ("lib", ".dylib"),
_ => ("lib", ".so"),
}
}
}
#[derive(Debug, Clone)]
pub struct WasmImport {
pub module: String,
pub name: String,
pub signature: String,
}
#[derive(Debug, Clone)]
pub struct WasmMemoryConfig {
pub initial_pages: usize,
pub max_pages: Option<usize>,
pub allow_growth: bool,
}
#[derive(Debug, Clone, PartialEq)]
pub enum PackageFormat {
Native,
WebAssembly,
CSharedLibrary,
AndroidAAR,
IOSFramework,
PythonWheel,
Docker,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GpuRequirements {
pub min_memory_mb: usize,
pub compute_capability: Option<String>,
pub drivers: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MobileArchitecture {
ARM64,
X86_64,
Universal,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CpuRequirements {
pub min_cores: usize,
pub instruction_sets: Vec<String>,
pub min_frequency_mhz: Option<usize>,
}
pub struct ModelServer<F: Float + Debug + scirs2_core::ndarray::ScalarOperand + NumAssign + 'static>
{
pub(super) model: Sequential<F>,
pub(super) config: ServerConfig,
pub(super) stats: ServerStats,
}
impl<
F: Float
+ Debug
+ scirs2_core::ndarray::ScalarOperand
+ FromPrimitive
+ Display
+ NumAssign
+ 'static,
> ModelServer<F>
{
pub fn new(model: Sequential<F>, config: ServerConfig) -> Self {
Self {
model,
config,
stats: ServerStats {
total_requests: 0,
successful_predictions: 0,
total_errors: 0,
avg_response_time_ms: 0.0,
active_requests: 0,
},
}
}
pub fn start(&mut self) -> Result<()> {
println!("Starting SciRS2 Model Server on port {}", self.config.port);
println!("Max batch size: {}", self.config.max_batch_size);
println!("Timeout: {}s", self.config.timeout_seconds);
Ok(())
}
pub fn predict(&mut self, input: &ArrayD<F>) -> Result<ArrayD<F>> {
self.stats.total_requests += 1;
self.stats.active_requests += 1;
let start_time = std::time::Instant::now();
let result = self.model.forward(input);
let elapsed = start_time.elapsed();
self.stats.active_requests -= 1;
match result {
Ok(output) => {
self.stats.successful_predictions += 1;
self.update_response_time(elapsed.as_millis() as f64);
Ok(output)
}
Err(e) => {
self.stats.total_errors += 1;
Err(e)
}
}
}
pub fn get_stats(&self) -> &ServerStats {
&self.stats
}
pub(super) fn update_response_time(&mut self, response_time_ms: f64) {
let total_responses = self.stats.successful_predictions + self.stats.total_errors;
if total_responses > 0 {
self.stats.avg_response_time_ms =
(self.stats.avg_response_time_ms * (total_responses - 1) as f64 + response_time_ms)
/ total_responses as f64;
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ServerStats {
pub total_requests: u64,
pub successful_predictions: u64,
pub total_errors: u64,
pub avg_response_time_ms: f64,
pub active_requests: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MobilePlatform {
IOS,
Android,
Universal,
}
#[derive(Debug, Clone)]
pub struct FrameworkConfig {
pub use_metal: bool,
pub use_nnapi: bool,
pub use_gpu: bool,
pub thread_pool_size: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PackageMetadata {
pub name: String,
pub version: String,
pub description: String,
pub author: String,
pub license: String,
pub platforms: Vec<String>,
pub dependencies: HashMap<String, String>,
pub input_specs: Vec<TensorSpec>,
pub output_specs: Vec<TensorSpec>,
pub runtime_requirements: RuntimeRequirements,
pub timestamp: String,
pub checksum: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TensorSpec {
pub name: String,
pub shape: Vec<Option<usize>>,
pub dtype: String,
pub description: Option<String>,
pub range: Option<(f64, f64)>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CallingConvention {
CDecl,
StdCall,
FastCall,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RuntimeRequirements {
pub min_memory_mb: usize,
pub cpu_requirements: CpuRequirements,
pub gpu_requirements: Option<GpuRequirements>,
pub system_dependencies: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct WasmConfig {
pub wasm_version: WasmVersion,
pub enable_simd: bool,
pub enable_threads: bool,
pub memory_config: WasmMemoryConfig,
pub imports: Vec<WasmImport>,
pub exports: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct CBindingConfig {
pub library_name: String,
pub header_guard: String,
pub namespace: Option<String>,
pub calling_convention: CallingConvention,
pub additional_headers: Vec<String>,
pub type_mappings: HashMap<String, String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WasmVersion {
V1_0,
V2_0,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OptimizationLevel {
None,
Basic,
Aggressive,
Size,
}