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

1""" 

2Registry-based converter infrastructure using metaclass-registry. 

3 

4This module provides the ConverterBase class using AutoRegisterMeta, 

5concrete converter implementations for each framework, and a helper 

6function for registry lookups. 

7""" 

8 

9from abc import abstractmethod 

10from collections.abc import Mapping 

11from types import MappingProxyType 

12from typing import ClassVar 

13 

14from metaclass_registry import AutoRegisterMeta 

15 

16from arraybridge.types import MemoryType 

17 

18 

19class ConverterBase(metaclass=AutoRegisterMeta): 

20 """Base class for memory type converters using auto-registration. 

21 

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 """ 

25 

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 

30 

31 @abstractmethod 

32 def to_numpy(self, data, gpu_id): 

33 """Extract to NumPy (type-specific implementation).""" 

34 pass 

35 

36 @abstractmethod 

37 def from_numpy(self, data, gpu_id): 

38 """Create from NumPy (type-specific implementation).""" 

39 pass 

40 

41 @abstractmethod 

42 def from_dlpack(self, data, gpu_id): 

43 """Create from DLPack capsule (type-specific implementation).""" 

44 pass 

45 

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 

50 

51 

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) 

56 

57 to_numpy.__qualname__ = f"{mem_type.value.capitalize()}Converter.to_numpy" 

58 return to_numpy 

59 

60 

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) 

65 

66 from_numpy.__qualname__ = f"{mem_type.value.capitalize()}Converter.from_numpy" 

67 return from_numpy 

68 

69 

70def _make_device_mover(mem_type: MemoryType): 

71 """Create a converter adapter over declaration-owned device movement.""" 

72 

73 def move_to_device(self, data, gpu_id): 

74 del self 

75 return mem_type.move_to_device(data, gpu_id) 

76 

77 move_to_device.__qualname__ = f"{mem_type.value.capitalize()}Converter.move_to_device" 

78 return move_to_device 

79 

80 

81def _make_dlpack_importer(mem_type: MemoryType): 

82 """Create a converter adapter over declaration-owned DLPack import.""" 

83 

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) 

89 

90 from_dlpack.__qualname__ = f"{mem_type.value.capitalize()}Converter.from_dlpack" 

91 return from_dlpack 

92 

93 

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) 

107 

108 

109# Create all converter classes at module load time 

110_create_converter_classes() 

111 

112 

113def get_converter(memory_type: str): 

114 """Get a converter instance for the given memory type. 

115 

116 Args: 

117 memory_type: The memory type string (e.g., "numpy", "torch") 

118 

119 Returns: 

120 A converter instance for the memory type 

121 

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() 

132 

133 

134def _add_converter_methods(): 

135 """Add to_X() methods to ConverterBase. 

136 

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}" 

142 

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) 

147 

148 return method 

149 

150 setattr(ConverterBase, method_name, make_method(target_type)) 

151 

152 

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()) 

157 

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)}") 

167 

168 

169# Add to_X() conversion methods after converter classes are created 

170_add_converter_methods() 

171 

172# Run validation at module load time 

173_validate_registry() 

174ConverterBase.__registry__ = MappingProxyType(dict(ConverterBase.__registry__))