Coverage for src/arraybridge/converters_registry.py: 81%
83 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 11:15 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 11:15 +0000
1"""
2Registry-based converter infrastructure using metaclass-registry.
4This module provides the ConverterBase class using AutoRegisterMeta,
5concrete converter implementations for each framework, and a helper
6function for registry lookups.
7"""
9from abc import abstractmethod
10from collections.abc import Mapping
11from types import MappingProxyType
12from typing import ClassVar
14from metaclass_registry import AutoRegisterMeta
16from arraybridge.types import MemoryType
19class ConverterBase(metaclass=AutoRegisterMeta):
20 """Base class for memory type converters using auto-registration.
22 Each concrete converter sets memory_type to register itself in the registry.
23 The registry key is the memory_type attribute (e.g., "numpy", "torch").
24 """
26 __registry_key__ = "memory_type"
27 # Simple dict: converters are created in this module, without lazy discovery.
28 __registry__: ClassVar[Mapping[str, type["ConverterBase"]]] = {}
29 memory_type: str | None = None
31 @abstractmethod
32 def to_numpy(self, data, gpu_id):
33 """Extract to NumPy (type-specific implementation)."""
34 pass
36 @abstractmethod
37 def from_numpy(self, data, gpu_id):
38 """Create from NumPy (type-specific implementation)."""
39 pass
41 @abstractmethod
42 def from_dlpack(self, data, gpu_id):
43 """Create from DLPack capsule (type-specific implementation)."""
44 pass
46 @abstractmethod
47 def move_to_device(self, data, gpu_id):
48 """Move data to specified GPU device if needed (type-specific implementation)."""
49 pass
52def _make_to_numpy(mem_type: MemoryType):
53 def to_numpy(self, data, gpu_id):
54 del self, gpu_id
55 return mem_type.to_numpy(data)
57 to_numpy.__qualname__ = f"{mem_type.value.capitalize()}Converter.to_numpy"
58 return to_numpy
61def _make_from_numpy(mem_type: MemoryType):
62 def from_numpy(self, data, gpu_id):
63 del self
64 return mem_type.from_numpy(data, gpu_id)
66 from_numpy.__qualname__ = f"{mem_type.value.capitalize()}Converter.from_numpy"
67 return from_numpy
70def _make_device_mover(mem_type: MemoryType):
71 """Create a converter adapter over declaration-owned device movement."""
73 def move_to_device(self, data, gpu_id):
74 del self
75 return mem_type.move_to_device(data, gpu_id)
77 move_to_device.__qualname__ = f"{mem_type.value.capitalize()}Converter.move_to_device"
78 return move_to_device
81def _make_dlpack_importer(mem_type: MemoryType):
82 """Create a converter adapter over declaration-owned DLPack import."""
84 def from_dlpack(self, data, gpu_id):
85 del self
86 module = mem_type.import_module()
87 with mem_type.device_scope(gpu_id, module):
88 return mem_type.from_dlpack(data, module)
90 from_dlpack.__qualname__ = f"{mem_type.value.capitalize()}Converter.from_dlpack"
91 return from_dlpack
94# Auto-generate converter classes for each memory type
95def _create_converter_classes():
96 """Create concrete converter classes for each memory type."""
97 for mem_type in MemoryType:
98 class_attrs = {
99 "memory_type": mem_type.value,
100 "to_numpy": _make_to_numpy(mem_type),
101 "from_numpy": _make_from_numpy(mem_type),
102 "move_to_device": _make_device_mover(mem_type),
103 "from_dlpack": _make_dlpack_importer(mem_type),
104 }
105 class_name = f"{mem_type.value.capitalize()}Converter"
106 type(class_name, (ConverterBase,), class_attrs)
109# Create all converter classes at module load time
110_create_converter_classes()
113def get_converter(memory_type: str):
114 """Get a converter instance for the given memory type.
116 Args:
117 memory_type: The memory type string (e.g., "numpy", "torch")
119 Returns:
120 A converter instance for the memory type
122 Raises:
123 ValueError: If memory type is not registered
124 """
125 converter_class = ConverterBase.__registry__.get(memory_type)
126 if converter_class is None:
127 raise ValueError(
128 f"No converter registered for memory type '{memory_type}'. "
129 f"Available types: {sorted(ConverterBase.__registry__.keys())}"
130 )
131 return converter_class()
134def _add_converter_methods():
135 """Add to_X() methods to ConverterBase.
137 For each target memory type, generates a method like to_cupy(), to_torch(), etc.
138 that tries GPU-to-GPU conversion via DLPack first, then falls back to CPU roundtrip.
139 """
140 for target_type in MemoryType:
141 method_name = f"to_{target_type.value}"
143 def make_method(tgt):
144 def method(self, data, gpu_id):
145 source_type = MemoryType(self.memory_type)
146 return source_type.convert_to(data, tgt, gpu_id)
148 return method
150 setattr(ConverterBase, method_name, make_method(target_type))
153def _validate_registry():
154 """Validate that all memory types are registered."""
155 required_types = {mt.value for mt in MemoryType}
156 registered_types = set(ConverterBase.__registry__.keys())
158 if required_types != registered_types:
159 missing = required_types - registered_types
160 extra = registered_types - required_types
161 msg_parts = []
162 if missing:
163 msg_parts.append(f"Missing: {missing}")
164 if extra:
165 msg_parts.append(f"Extra: {extra}")
166 raise RuntimeError(f"Registry validation failed. {', '.join(msg_parts)}")
169# Add to_X() conversion methods after converter classes are created
170_add_converter_methods()
172# Run validation at module load time
173_validate_registry()
174ConverterBase.__registry__ = MappingProxyType(dict(ConverterBase.__registry__))