Coverage for gamdpy/interactions/interaction.py: 44%

121 statements  

« prev     ^ index     » next       coverage.py v7.9.1, created at 2025-06-14 15:55 +0200

1import numba 

2from numba import cuda 

3from abc import ABC, abstractmethod 

4from gamdpy import Configuration 

5from typing import Callable 

6 

7class Interaction(ABC): 

8 """ 

9 Abstract Base Class specifying the requirements for an interaction 

10 """ 

11 

12 @abstractmethod 

13 def get_kernel(self, configuration: Configuration, compute_plan: dict, compute_flags: dict[str,bool]) -> Callable: 

14 """ 

15 Get a kernel (or python function depending on compute_plan["gridsync"]) that implements calculation of the interaction 

16 """ 

17 

18 pass 

19 

20 @abstractmethod 

21 def get_params(self, configuration: Configuration, compute_plan: dict) -> tuple: 

22 """ 

23 Get a tuple with the parameters expected by the associated kernel 

24 """ 

25 

26 pass 

27 

28 def check_datastructure_validity(self) -> bool: 

29 """ 

30 Interactions which have an internal data structure (think: PairPotential and its NbList) should overwrite this method  

31 with one that checks the validity of it, and throws an error if not valid (Later: repair it and signal rerun of timeblock)  

32 """ 

33 return True 

34 

35def merge_interactions(configuration: Configuration, kernelA: Callable, paramsA: tuple, interactionB: Interaction, compute_plan: dict, compute_flags: dict[str,bool]) -> tuple[Callable, tuple] : 

36 paramsB = interactionB.get_params(configuration, compute_plan) 

37 kernelB = interactionB.get_kernel(configuration, compute_plan, compute_flags) 

38 

39 if compute_plan['gridsync']: 

40 # A device function, calling a number of device functions, using gridsync to syncronize 

41 @cuda.jit( device=compute_plan['gridsync']) 

42 def interactions(grid, vectors, scalars, ptype, sim_box, interaction_parameters): 

43 kernelA(grid, vectors, scalars, ptype, sim_box, interaction_parameters[0]) 

44 grid.sync() # Not always necessary !!! 

45 kernelB(grid, vectors, scalars, ptype, sim_box, interaction_parameters[1]) 

46 return 

47 return interactions, (paramsA, paramsB, ) 

48 else: 

49 # A python function, making several kernel calls to syncronize  

50 def interactions(grid, vectors, scalars, ptype, sim_box, interaction_parameters): 

51 kernelA(0, vectors, scalars, ptype, sim_box, interaction_parameters[0]) 

52 kernelB(0, vectors, scalars, ptype, sim_box, interaction_parameters[1]) 

53 return 

54 return interactions, (paramsA, paramsB, ) 

55 

56 

57def add_interactions_list(configuration: Configuration, interactions_list: list[Interaction], compute_plan: dict, compute_flags: dict[str,bool], verbose: bool = False) -> tuple[Callable, tuple]: 

58 

59 # Setup first interaction and cuda.jit it if gridsync is used for syncronization 

60 params = get_initializer_params(configuration, compute_plan) 

61 kernel: Callable = get_initializer_kernel(configuration, compute_plan, compute_flags) 

62 if compute_plan['gridsync']: 

63 kernel: Callable = cuda.jit( device=compute_plan['gridsync'] )(kernel) 

64 

65 # Merge in the rest of the interaction (maximum recursion depth might set a maximum for number of interactions) 

66 for i in range(len(interactions_list)): 

67 kernel, params = merge_interactions(configuration, kernel, params, interactions_list[i], compute_plan, compute_flags) 

68 

69 return kernel, params 

70 

71 

72def get_initializer_params(configuration, compute_plan): 

73 return (0,) 

74 

75 

76def get_initializer_kernel(configuration, compute_plan, compute_flags) -> Callable: 

77 

78 num_cscalars = configuration.num_cscalars 

79 compute_stresses = compute_flags['stresses'] 

80 

81 # Unpack parameters from configuration and compute_plan 

82 D, num_part = configuration.D, configuration.N 

83 pb, tp, gridsync, UtilizeNIII = [compute_plan[key] for key in ['pb', 'tp', 'gridsync', 'UtilizeNIII']] 

84 num_blocks = (num_part - 1) // pb + 1 

85 

86 # Unpack indices for vectors and scalars to be compiled into kernel 

87 f_id, = [configuration.vectors.indices[key] for key in ['f']] 

88 

89 if compute_stresses: 

90 sx_id = configuration.vectors.indices['sx'] 

91 if D > 1: 

92 sy_id = configuration.vectors.indices['sy'] 

93 if D > 2: 

94 sz_id = configuration.vectors.indices['sz'] 

95 if D > 3: 

96 sw_id = configuration.vectors.indices['sw'] 

97 

98 

99 def kernel(grid, vectors, scalars, ptype, sim_box, interaction_parameters): 

100 

101 global_id, my_t = cuda.grid(2) 

102 

103 if global_id < num_part and my_t==0: 

104 for k in range(num_cscalars): 

105 scalars[global_id, k] = numba.float32(0.0) 

106 

107 

108 if global_id < num_part and my_t==0: # Initializion of forces moved here to make NewtonIII possible  

109 for k in range(D): 

110 vectors[f_id][global_id, k] = numba.float32(0.0) 

111 if compute_stresses: 

112 vectors[sx_id][global_id, k] = numba.float32(0.0) 

113 if D > 1: 

114 vectors[sy_id][global_id, k] = numba.float32(0.0) 

115 if D > 2: 

116 vectors[sz_id][global_id, k] = numba.float32(0.0) 

117 if D > 3: 

118 vectors[sw_id][global_id, k] = numba.float32(0.0) 

119 return 

120 

121 

122 if gridsync: 

123 # A device function, calling a number of device functions, using gridsync to syncronize 

124 return cuda.jit( device=gridsync )(kernel) 

125 else: 

126 return cuda.jit( device=gridsync )(kernel)[num_blocks, (pb, tp)] 

127 

128 

129 

130 

131# Function below not used and will be removed 

132 

133def add_interactions_list_old(configuration, interactions_list, compute_plan, compute_flags, verbose=True,): 

134 gridsync = compute_plan['gridsync'] 

135 num_interactions = len(interactions_list) 

136 assert 0 < num_interactions <= 5 

137 

138 interaction_params_list = [] 

139 for interaction in interactions_list: 

140 interaction_params_list.append(interaction.get_params(configuration, compute_plan, verbose=verbose)) 

141 

142 i0 = interactions_list[0].get_kernel(configuration, compute_plan, compute_flags, verbose=verbose) 

143 if num_interactions>1: 

144 i1 = interactions_list[1].get_kernel(configuration, compute_plan, compute_flags, verbose=verbose) 

145 if num_interactions>2: 

146 i2 = interactions_list[2].get_kernel(configuration, compute_plan, compute_flags, verbose=verbose) 

147 if num_interactions>3: 

148 i3 = interactions_list[3].get_kernel(configuration, compute_plan, compute_flags, verbose=verbose) 

149 if num_interactions>4: 

150 i4 = interactions_list[4].get_kernel(configuration, compute_plan, compute_flags, verbose=verbose) 

151 

152 if gridsync: 

153 # A device function, calling a number of device functions, using gridsync to syncronize 

154 @cuda.jit( device=gridsync ) 

155 def interactions(grid, vectors, scalars, ptype, sim_box, interaction_parameters): 

156 i0(grid, vectors, scalars, ptype, sim_box, interaction_parameters[0]) 

157 if num_interactions>1: 

158 grid.sync() # Not always necessary !!! 

159 i1(grid, vectors, scalars, ptype, sim_box, interaction_parameters[1]) 

160 if num_interactions>2: 

161 grid.sync() # Not always necessary !!! 

162 i2(grid, vectors, scalars, ptype, sim_box, interaction_parameters[2]) 

163 if num_interactions>3: 

164 grid.sync() # Not always necessary !!! 

165 i3(grid, vectors, scalars, ptype, sim_box, interaction_parameters[3]) 

166 if num_interactions>4: 

167 grid.sync() # Not always necessary !!! 

168 i4(grid, vectors, scalars, ptype, sim_box, interaction_parameters[4]) 

169 return 

170 return interactions, tuple(interaction_params_list) 

171 

172 else: 

173 # A python function, making several kernel calls to syncronize  

174 #@cuda.jit( device=gridsync ) 

175 def interactions(grid, vectors, scalars, ptype, sim_box, interaction_parameters): 

176 i0(0, vectors, scalars, ptype, sim_box, interaction_parameters[0]) 

177 if num_interactions>1: 

178 i1(0, vectors, scalars, ptype, sim_box, interaction_parameters[1]) 

179 if num_interactions>2: 

180 i2(0, vectors, scalars, ptype, sim_box, interaction_parameters[2]) 

181 if num_interactions>3: 

182 i3(0, vectors, scalars, ptype, sim_box, interaction_parameters[3]) 

183 if num_interactions>4: 

184 i4(0, vectors, scalars, ptype, sim_box, interaction_parameters[4]) 

185 return 

186 return interactions, tuple(interaction_params_list)