(
mode: Union[
Literal["DirectAnswer"],
Literal["FullConnected"],
Literal["Random"],
Literal["Chain"],
Literal["Debate"],
Literal["Layered"],
Literal["Star"],
],
N: int,
)
| 212 | metrics_recorder.log_cumulative(batch_index=i_batch) |
| 213 | |
| 214 | def get_kwargs( |
| 215 | mode: Union[ |
| 216 | Literal["DirectAnswer"], |
| 217 | Literal["FullConnected"], |
| 218 | Literal["Random"], |
| 219 | Literal["Chain"], |
| 220 | Literal["Debate"], |
| 221 | Literal["Layered"], |
| 222 | Literal["Star"], |
| 223 | ], |
| 224 | N: int, |
| 225 | ): |
| 226 | fixed_spatial_masks: List[List[int]] = None |
| 227 | fixed_temporal_masks: List[List[int]] = None |
| 228 | node_kwargs = None |
| 229 | |
| 230 | def generate_layered_graph(n, layer_num=2): |
| 231 | adj_matrix = [[0 for _ in range(n)] for _ in range(n)] |
| 232 | base_size = n // layer_num |
| 233 | remainder = n % layer_num |
| 234 | layers: List[int] = [] |
| 235 | for i in range(layer_num): |
| 236 | size = base_size + (1 if i < remainder else 0) |
| 237 | layers.extend([i] * size) |
| 238 | random.shuffle(layers) |
| 239 | for i in range(n): |
| 240 | current_layer = layers[i] |
| 241 | for j in range(n): |
| 242 | if layers[j] == current_layer + 1: |
| 243 | adj_matrix[i][j] = 1 |
| 244 | return adj_matrix |
| 245 | |
| 246 | def generate_star_graph(n): |
| 247 | matrix = [[0] * n for _ in range(n)] |
| 248 | for i in range(n): |
| 249 | for j in range(i + 1, n): |
| 250 | matrix[i][j] = 1 |
| 251 | return matrix |
| 252 | |
| 253 | if mode == "DirectAnswer": |
| 254 | fixed_spatial_masks = [[0]] |
| 255 | fixed_temporal_masks = [[0]] |
| 256 | node_kwargs = [{"role": "Normal Programmer"}] |
| 257 | elif mode == "FullConnected": |
| 258 | fixed_spatial_masks = [[1 if i != j else 0 for i in range(N)] for j in range(N)] |
| 259 | fixed_temporal_masks = [[1 for _ in range(N)] for _ in range(N)] |
| 260 | elif mode == "Random": |
| 261 | fixed_spatial_masks = [[random.randint(0, 1) if i != j else 0 for i in range(N)] for j in range(N)] |
| 262 | fixed_temporal_masks = [[random.randint(0, 1) for _ in range(N)] for _ in range(N)] |
| 263 | elif mode == "Chain": |
| 264 | fixed_spatial_masks = [[1 if i == j + 1 else 0 for i in range(N)] for j in range(N)] |
| 265 | fixed_temporal_masks = [[1 if i == 0 and j == N - 1 else 0 for i in range(N)] for j in range(N)] |
| 266 | elif mode == "Debate": |
| 267 | fixed_spatial_masks = [[0 for _ in range(N)] for _ in range(N)] |
| 268 | fixed_temporal_masks = [[1 for _ in range(N)] for _ in range(N)] |
| 269 | elif mode == "Layered": |
| 270 | fixed_spatial_masks = generate_layered_graph(N) |
| 271 | fixed_temporal_masks = [[1 for _ in range(N)] for _ in range(N)] |
no test coverage detected