-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathalgo_registry.py
More file actions
31 lines (24 loc) · 787 Bytes
/
Copy pathalgo_registry.py
File metadata and controls
31 lines (24 loc) · 787 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
"""
Algorithm registry (open-source mirror).
ROUTING_REGISTRY: differentiable input routing (Sinkhorn v2 only).
"""
from typing import Dict, Type
from algo.routing.base import RouterBase
from algo.routing.sinkhorn_v2 import SinkhornRouterV2
ROUTING_REGISTRY: Dict[str, Type[RouterBase]] = {
"sinkhorn_v2": SinkhornRouterV2,
}
def build_router(name: str, **kwargs) -> RouterBase:
"""
Build a routing module by name.
Args:
name: registry key (sinkhorn_v2)
**kwargs: forwarded to the router constructor
Returns:
RouterBase subclass instance
"""
assert name in ROUTING_REGISTRY, (
f"Unknown routing algorithm: {name}. "
f"Available: {list(ROUTING_REGISTRY.keys())}"
)
return ROUTING_REGISTRY[name](**kwargs)