(
mode: Union[
Literal["DirectAnswer"],
Literal["FullConnected"],
Literal["Random"],
Literal["Chain"],
Literal["Debate"],
Literal["Layered"],
Literal["Star"],
Literal["Mesh"],
],
N: int,
)
| 178 | |
| 179 | |
| 180 | def get_kwargs( |
| 181 | mode: Union[ |
| 182 | Literal["DirectAnswer"], |
| 183 | Literal["FullConnected"], |
| 184 | Literal["Random"], |
| 185 | Literal["Chain"], |
| 186 | Literal["Debate"], |
| 187 | Literal["Layered"], |
| 188 | Literal["Star"], |
| 189 | Literal["Mesh"], |
| 190 | ], |
| 191 | N: int, |
| 192 | ): |
| 193 | fixed_spatial_masks: List[List[int]] = None |
| 194 | fixed_temporal_masks: List[List[int]] = None |
| 195 | node_kwargs = None |
| 196 | |
| 197 | def generate_layered_graph(n, layer_num=2): |
| 198 | adj_matrix = [[0] * n for _ in range(n)] |
| 199 | base_size = n // layer_num |
| 200 | remainder = n % layer_num |
| 201 | layers: List[int] = [] |
| 202 | for i in range(layer_num): |
| 203 | size = base_size + (1 if i < remainder else 0) |
| 204 | layers.extend([i] * size) |
| 205 | random.shuffle(layers) |
| 206 | for i in range(n): |
| 207 | current_layer = layers[i] |
| 208 | for j in range(n): |
| 209 | if layers[j] == current_layer + 1: |
| 210 | adj_matrix[i][j] = 1 |
| 211 | return adj_matrix |
| 212 | |
| 213 | def generate_mesh_graph(n): |
| 214 | adj_matrix = [[0] * n for _ in range(n)] |
| 215 | for i in range(n): |
| 216 | for j in range(i + 1, n): |
| 217 | adj_matrix[i][j] = 1 |
| 218 | return adj_matrix |
| 219 | |
| 220 | def generate_star_graph(n): |
| 221 | adj_matrix = [[0] * n for _ in range(n)] |
| 222 | for i in range(1, n): |
| 223 | adj_matrix[0][i] = 1 |
| 224 | return adj_matrix |
| 225 | |
| 226 | if mode == "DirectAnswer": |
| 227 | fixed_spatial_masks = [[0]] |
| 228 | fixed_temporal_masks = [[0]] |
| 229 | node_kwargs = [{"role": "Normal"}] |
| 230 | elif mode == "FullConnected": |
| 231 | fixed_spatial_masks = [[1 if i != j else 0 for i in range(N)] for j in range(N)] |
| 232 | fixed_temporal_masks = [[1 for _ in range(N)] for _ in range(N)] |
| 233 | elif mode == "Random": |
| 234 | fixed_spatial_masks = [[random.randint(0, 1) if i != j else 0 for i in range(N)] for j in range(N)] |
| 235 | fixed_temporal_masks = [[random.randint(0, 1) for _ in range(N)] for _ in range(N)] |
| 236 | elif mode == "Chain": |
| 237 | fixed_spatial_masks = [[1 if i == j + 1 else 0 for i in range(N)] for j in range(N)] |
no test coverage detected