use std::num::NonZeroUsize;
use eredu_checkpoint::{AffineQuantization, WeightQuantization};
use eredu_core::{
CompletionCancellationMode, DraftingPlan, ParallelRankTopology, PreparationPolicy,
QuantizationRequest, ResidencyRequest, SessionCapabilities,
};
use crate::{
CacheResidencyPolicy, CommunicationCompletionPolicy, LayerWeightResidency,
PipelineWireContract, WeightResidency,
};
#[derive(Debug, Clone, Copy, Default, Eq, PartialEq)]
pub enum DraftingLoadRequest {
#[default]
ArchitectureDefault,
Disabled,
Embedded {
max_draft_tokens: NonZeroUsize,
},
ExternalTarget,
}
impl DraftingLoadRequest {
pub fn embedded(max_draft_tokens: usize) -> Result<Self, NormalizedLoadRequestError> {
let max_draft_tokens = NonZeroUsize::new(max_draft_tokens)
.ok_or(NormalizedLoadRequestError::ZeroEmbeddedDraftCapacity)?;
Ok(Self::Embedded { max_draft_tokens })
}
pub const fn embedded_capacity(self) -> Option<NonZeroUsize> {
match self {
Self::Embedded { max_draft_tokens } => Some(max_draft_tokens),
Self::ArchitectureDefault | Self::Disabled | Self::ExternalTarget => None,
}
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct ParallelLoadRequest {
rank: ParallelRankTopology,
wire: PipelineWireContract,
maximum_batch_size: i32,
maximum_sequence_length: i32,
completion: CommunicationCompletionPolicy,
}
impl ParallelLoadRequest {
pub fn new(
rank: ParallelRankTopology,
wire: PipelineWireContract,
maximum_batch_size: i32,
maximum_sequence_length: i32,
completion: CommunicationCompletionPolicy,
) -> Result<Self, NormalizedLoadRequestError> {
if rank.is_replicated() {
return Err(NormalizedLoadRequestError::ReplicatedParallelTopology);
}
if maximum_batch_size <= 0 || maximum_sequence_length <= 0 {
return Err(NormalizedLoadRequestError::InvalidInvocationLimits {
maximum_batch_size,
maximum_sequence_length,
});
}
Ok(Self {
rank,
wire,
maximum_batch_size,
maximum_sequence_length,
completion,
})
}
pub const fn rank(self) -> ParallelRankTopology {
self.rank
}
pub const fn wire(self) -> PipelineWireContract {
self.wire
}
pub const fn invocation_limits(self) -> (i32, i32) {
(self.maximum_batch_size, self.maximum_sequence_length)
}
pub const fn completion(self) -> CommunicationCompletionPolicy {
self.completion
}
const fn with_completion(mut self, completion: CommunicationCompletionPolicy) -> Self {
self.completion = completion;
self
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum NormalizedLoadRequestError {
#[error("source reader-cache limit must be positive")]
ZeroCachedShards,
#[error(
"partitioned invocation limits must be positive, got batch {maximum_batch_size} and sequence {maximum_sequence_length}"
)]
InvalidInvocationLimits {
maximum_batch_size: i32,
maximum_sequence_length: i32,
},
#[error("explicit parallel execution requires a non-replicated topology")]
ReplicatedParallelTopology,
#[error("parallel execution cannot replace an existing local completion policy")]
ConflictingCompletionPolicy,
#[error("a communication completion policy requires a parallel topology")]
OrphanedModelCompletion,
#[error("embedded draft capacity must be positive")]
ZeroEmbeddedDraftCapacity,
#[error("unsupported speculative drafting plan")]
UnsupportedDraftingPlan,
#[error("{0}")]
Completion(String),
#[error("{0}")]
Quantization(String),
}
#[derive(Debug, Clone, Default, Eq, PartialEq)]
pub struct NormalizedLoadRequest {
quantization: Option<QuantizationRequest>,
max_cached_shards: Option<NonZeroUsize>,
parallel: Option<ParallelLoadRequest>,
communication_completion: Option<CommunicationCompletionPolicy>,
weight_residency: WeightResidency,
state_residency: CacheResidencyPolicy,
required_session_capabilities: SessionCapabilities,
prompt_cache_persistence: bool,
drafting: DraftingLoadRequest,
}
#[derive(Debug, Clone, Copy)]
pub struct ValidatedModelLoadRequest<'a> {
request: &'a NormalizedLoadRequest,
policy: PreparationPolicy,
}
impl<'a> ValidatedModelLoadRequest<'a> {
pub const fn preparation_policy(self) -> PreparationPolicy {
self.policy
}
pub const fn request(self) -> &'a NormalizedLoadRequest {
self.request
}
}
impl NormalizedLoadRequest {
pub const fn with_prompt_cache_persistence(mut self, required: bool) -> Self {
self.prompt_cache_persistence = required;
self
}
pub const fn prompt_cache_persistence(&self) -> bool {
self.prompt_cache_persistence
}
pub const fn with_max_cached_shards(mut self, maximum: NonZeroUsize) -> Self {
self.max_cached_shards = Some(maximum);
self
}
pub const fn max_cached_shards(&self) -> usize {
match self.max_cached_shards {
Some(maximum) => maximum.get(),
None => self.weight_residency.max_cached_shards(),
}
}
pub fn with_quantization(quantization: QuantizationRequest) -> Self {
Self {
quantization: Some(quantization),
..Self::default()
}
}
pub fn with_parallel_execution(
mut self,
parallel: ParallelLoadRequest,
) -> Result<Self, NormalizedLoadRequestError> {
if self.communication_completion.is_some() {
return Err(NormalizedLoadRequestError::ConflictingCompletionPolicy);
}
self.parallel = Some(parallel);
Ok(self)
}
pub const fn with_communication_completion_policy(
mut self,
policy: CommunicationCompletionPolicy,
) -> Self {
self.set_communication_completion_policy(policy);
self
}
pub const fn set_communication_completion_policy(
&mut self,
policy: CommunicationCompletionPolicy,
) {
match self.parallel {
Some(parallel) => self.parallel = Some(parallel.with_completion(policy)),
None => self.communication_completion = Some(policy),
}
}
pub fn with_weight_residency(mut self, residency: WeightResidency) -> Self {
self.weight_residency = residency;
self
}
pub fn with_state_residency(mut self, residency: CacheResidencyPolicy) -> Self {
self.state_residency = residency;
self
}
pub const fn with_required_session_capabilities(
mut self,
capabilities: SessionCapabilities,
) -> Self {
self.required_session_capabilities = capabilities;
self
}
pub const fn with_drafting(mut self, drafting: DraftingLoadRequest) -> Self {
self.drafting = drafting;
self
}
pub fn with_drafting_plan(
self,
plan: &DraftingPlan,
) -> Result<Self, NormalizedLoadRequestError> {
let drafting = match plan {
DraftingPlan::Disabled => DraftingLoadRequest::Disabled,
DraftingPlan::Embedded {
max_draft_tokens, ..
} => DraftingLoadRequest::embedded(*max_draft_tokens)?,
DraftingPlan::External { .. } => DraftingLoadRequest::ExternalTarget,
_ => return Err(NormalizedLoadRequestError::UnsupportedDraftingPlan),
};
Ok(self.with_drafting(drafting))
}
pub const fn quantization(&self) -> Option<QuantizationRequest> {
self.quantization
}
pub const fn parallel_execution(&self) -> Option<ParallelLoadRequest> {
self.parallel
}
pub const fn parallel_topology(&self) -> Option<ParallelRankTopology> {
match self.parallel {
Some(parallel) => Some(parallel.rank()),
None => None,
}
}
pub const fn pipeline_wire_contract(&self) -> Option<PipelineWireContract> {
match self.parallel {
Some(parallel) => Some(parallel.wire()),
None => None,
}
}
pub const fn has_parallel_execution(&self) -> bool {
self.parallel.is_some()
}
pub fn partitioned_invocation_limits(&self) -> Option<(i32, i32)> {
self.parallel.map(ParallelLoadRequest::invocation_limits)
}
pub fn communication_completion_policy(
&self,
) -> Result<Option<CommunicationCompletionPolicy>, NormalizedLoadRequestError> {
match (self.parallel, self.communication_completion) {
(Some(parallel), None) => Ok(Some(parallel.completion())),
(None, None) => Ok(None),
(None, Some(_)) => Err(NormalizedLoadRequestError::OrphanedModelCompletion),
(Some(_), Some(_)) => Err(NormalizedLoadRequestError::ConflictingCompletionPolicy),
}
}
pub fn realtime_completion_policy(
&self,
) -> Result<CommunicationCompletionPolicy, NormalizedLoadRequestError> {
if let Some(parallel) = self.parallel {
return Ok(parallel.completion());
}
self.communication_completion.map_or_else(
|| {
CommunicationCompletionPolicy::new(
std::time::Duration::from_secs(30),
CompletionCancellationMode::QuarantineUntilComplete,
)
.map_err(|error| NormalizedLoadRequestError::Completion(error.to_string()))
},
Ok,
)
}
pub const fn weight_residency(&self) -> WeightResidency {
self.weight_residency
}
pub const fn state_residency(&self) -> &CacheResidencyPolicy {
&self.state_residency
}
pub const fn required_session_capabilities(&self) -> SessionCapabilities {
self.required_session_capabilities
}
pub const fn drafting(&self) -> DraftingLoadRequest {
self.drafting
}
pub fn weight_quantization(
&self,
) -> Result<Option<WeightQuantization>, NormalizedLoadRequestError> {
self.quantization
.map(|request| match request {
QuantizationRequest::Affine { group_size, bits } => {
let group_size = i32::try_from(group_size).map_err(|_| {
NormalizedLoadRequestError::Quantization(format!(
"group_size must fit in i32, got {group_size}"
))
})?;
AffineQuantization::new(group_size, i32::from(bits))
.map(WeightQuantization::Affine)
.map_err(|error| {
NormalizedLoadRequestError::Quantization(error.to_string())
})
}
QuantizationRequest::MxFp4 => Ok(WeightQuantization::MxFp4),
_ => Err(NormalizedLoadRequestError::Quantization(
"unknown load-time transformation request".into(),
)),
})
.transpose()
}
pub fn validate_model_preparation(
&self,
) -> Result<ValidatedModelLoadRequest<'_>, NormalizedLoadRequestError> {
if self.max_cached_shards() == 0 {
return Err(NormalizedLoadRequestError::ZeroCachedShards);
}
self.weight_quantization()?;
self.communication_completion_policy()?;
Ok(ValidatedModelLoadRequest {
request: self,
policy: self.project_preparation_policy(),
})
}
pub fn preparation_policy(&self) -> Result<PreparationPolicy, NormalizedLoadRequestError> {
self.validate_model_preparation()
.map(ValidatedModelLoadRequest::preparation_policy)
}
fn project_preparation_policy(&self) -> PreparationPolicy {
let residency = if self.weight_residency.parameter_bank_cache().is_some() {
ResidencyRequest::AddressableParameterBanks
} else {
match self.weight_residency.layers() {
LayerWeightResidency::FullyResident => ResidencyRequest::FullyResident,
LayerWeightResidency::LayerwiseHost(_) => ResidencyRequest::LayerwiseHost,
LayerWeightResidency::DenseDiskStream(_) => ResidencyRequest::DenseDiskStream,
}
};
let mut policy = PreparationPolicy::new(self.quantization, residency)
.with_required_session_capabilities(self.required_session_capabilities);
if let Some(topology) = self.parallel_topology() {
policy = policy.with_topology(topology.topology());
}
policy
}
}
#[cfg(test)]
mod tests {
use super::*;
use eredu_core::{ParallelTopology, QuantizationRequest};
fn completion() -> CommunicationCompletionPolicy {
CommunicationCompletionPolicy::new(
std::time::Duration::from_secs(1),
CompletionCancellationMode::QuarantineUntilComplete,
)
.unwrap()
}
#[test]
fn parallel_policy_is_atomic_and_exact() {
let rank =
ParallelRankTopology::new(ParallelTopology::new(2, 1, 1, 1).unwrap(), 1).unwrap();
let parallel = ParallelLoadRequest::new(
rank,
PipelineWireContract::new(crate::PipelineActivationDtype::Float32),
2,
128,
completion(),
)
.unwrap();
let request = NormalizedLoadRequest::with_quantization(QuantizationRequest::MxFp4)
.with_parallel_execution(parallel)
.unwrap()
.with_required_session_capabilities(SessionCapabilities::new(true, false, true));
request.validate_model_preparation().unwrap();
assert_eq!(request.parallel_execution(), Some(parallel));
assert_eq!(request.parallel_topology(), Some(rank));
assert_eq!(request.partitioned_invocation_limits(), Some((2, 128)));
assert_eq!(
request.communication_completion_policy().unwrap(),
Some(completion())
);
assert_eq!(
request.preparation_policy().unwrap().topology(),
Some(rank.topology())
);
}
#[test]
fn invalid_parallel_geometry_fails_at_parallel_construction() {
let rank =
ParallelRankTopology::new(ParallelTopology::new(2, 1, 1, 1).unwrap(), 0).unwrap();
let request = ParallelLoadRequest::new(
rank,
PipelineWireContract::new(crate::PipelineActivationDtype::Float32),
0,
128,
completion(),
);
assert!(matches!(
request,
Err(NormalizedLoadRequestError::InvalidInvocationLimits { .. })
));
let request = ParallelLoadRequest::new(
rank,
PipelineWireContract::new(crate::PipelineActivationDtype::Float32),
1,
-1,
completion(),
);
assert!(matches!(
request,
Err(NormalizedLoadRequestError::InvalidInvocationLimits { .. })
));
}
#[test]
fn replicated_topology_is_rejected_at_parallel_construction() {
let rank =
ParallelRankTopology::new(ParallelTopology::new(1, 1, 1, 1).unwrap(), 0).unwrap();
assert!(matches!(
ParallelLoadRequest::new(
rank,
PipelineWireContract::new(crate::PipelineActivationDtype::Float32),
1,
128,
completion(),
),
Err(NormalizedLoadRequestError::ReplicatedParallelTopology)
));
}
#[test]
fn local_completion_cannot_be_reinterpreted_as_parallel_completion() {
let rank =
ParallelRankTopology::new(ParallelTopology::new(2, 1, 1, 1).unwrap(), 0).unwrap();
let parallel = ParallelLoadRequest::new(
rank,
PipelineWireContract::new(crate::PipelineActivationDtype::Float32),
1,
128,
completion(),
)
.unwrap();
let request =
NormalizedLoadRequest::default().with_communication_completion_policy(completion());
assert!(matches!(
request.with_parallel_execution(parallel),
Err(NormalizedLoadRequestError::ConflictingCompletionPolicy)
));
}
#[test]
fn embedded_drafting_capacity_is_positive_by_construction() {
assert!(matches!(
DraftingLoadRequest::embedded(0),
Err(NormalizedLoadRequestError::ZeroEmbeddedDraftCapacity)
));
assert_eq!(
DraftingLoadRequest::embedded(4)
.unwrap()
.embedded_capacity()
.unwrap()
.get(),
4
);
}
#[test]
fn local_realtime_completion_does_not_become_valid_model_communication() {
let request =
NormalizedLoadRequest::default().with_communication_completion_policy(completion());
assert_eq!(request.realtime_completion_policy().unwrap(), completion());
assert!(matches!(
request.validate_model_preparation(),
Err(NormalizedLoadRequestError::OrphanedModelCompletion)
));
}
}