[Hybrid] Added supports_mamba_prefix_caching Protocol (#27339)
Signed-off-by: asafg <39553475+Josephasafg@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
parent
f4e8154076
commit
9273754222
@@ -697,6 +697,34 @@ def has_noops(
|
||||
return getattr(model, "has_noops", False)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SupportsMambaPrefixCaching(Protocol):
|
||||
"""The interface for models whose mamba layers support prefix caching.
|
||||
|
||||
This is currently experimental.
|
||||
"""
|
||||
|
||||
supports_mamba_prefix_caching: ClassVar[Literal[True]] = True
|
||||
|
||||
|
||||
@overload
|
||||
def supports_mamba_prefix_caching(
|
||||
model: object,
|
||||
) -> TypeIs[SupportsMambaPrefixCaching]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def supports_mamba_prefix_caching(
|
||||
model: type[object],
|
||||
) -> TypeIs[type[SupportsMambaPrefixCaching]]: ...
|
||||
|
||||
|
||||
def supports_mamba_prefix_caching(
|
||||
model: type[object] | object,
|
||||
) -> TypeIs[type[SupportsMambaPrefixCaching]] | TypeIs[SupportsMambaPrefixCaching]:
|
||||
return getattr(model, "supports_mamba_prefix_caching", False)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SupportsCrossEncoding(Protocol):
|
||||
"""The interface required for all models that support cross encoding."""
|
||||
|
||||
Reference in New Issue
Block a user