(
mode: Union[
Literal["DirectAnswer"],
Literal["FullConnected"],
Literal["Random"],
Literal["Chain"],
Literal["Debate"],
Literal["Layered"],
Literal["Star"],
Literal["Mesh"],
],
N: int,
)
| 134 | |
| 135 | |
| 136 | def get_kwargs( |
| 137 | mode: Union[ |
| 138 | Literal["DirectAnswer"], |
| 139 | Literal["FullConnected"], |
| 140 | Literal["Random"], |
| 141 | Literal["Chain"], |
| 142 | Literal["Debate"], |
| 143 | Literal["Layered"], |
| 144 | Literal["Star"], |
| 145 | Literal["Mesh"], |
| 146 | ], |
| 147 | N: int, |
| 148 | ): |
| 149 | fixed_spatial_masks: List[List[int]] = None |
| 150 | fixed_temporal_masks: List[List[int]] = None |
| 151 | node_kwargs = None |
| 152 | |
| 153 | def generate_layered_graph(n, layer_num=2): |
| 154 | adj_matrix = [[0] * n for _ in range(n)] |
| 155 | base_size = n // layer_num |
| 156 | remainder = n % layer_num |
| 157 | layers: List[int] = [] |
| 158 | for i in range(layer_num): |
| 159 | size = base_size + (1 if i < remainder else 0) |
| 160 | layers.extend([i] * size) |
| 161 | random.shuffle(layers) |
| 162 | for i in range(n): |
| 163 | current_layer = layers[i] |
| 164 | for j in range(n): |
| 165 | if layers[j] == current_layer + 1: |
| 166 | adj_matrix[i][j] = 1 |
| 167 | return adj_matrix |
| 168 | |
| 169 | def generate_mesh_graph(n): |
| 170 | adj_matrix = [[0] * n for _ in range(n)] |
| 171 | for i in range(n): |
| 172 | for j in range(i + 1, n): |
| 173 | adj_matrix[i][j] = 1 |
| 174 | return adj_matrix |
| 175 | |
| 176 | def generate_star_graph(n): |
| 177 | adj_matrix = [[0] * n for _ in range(n)] |
| 178 | for i in range(1, n): |
| 179 | adj_matrix[0][i] = 1 |
| 180 | return adj_matrix |
| 181 | |
| 182 | if mode == "DirectAnswer": |
| 183 | fixed_spatial_masks = [[0]] |
| 184 | fixed_temporal_masks = [[0]] |
| 185 | node_kwargs = [{"role": "Normal"}] |
| 186 | elif mode == "FullConnected": |
| 187 | fixed_spatial_masks = [[1 if i != j else 0 for i in range(N)] for j in range(N)] |
| 188 | fixed_temporal_masks = [[1 for _ in range(N)] for _ in range(N)] |
| 189 | elif mode == "Random": |
| 190 | fixed_spatial_masks = [[random.randint(0, 1) if i != j else 0 for i in range(N)] for j in range(N)] |
| 191 | fixed_temporal_masks = [[random.randint(0, 1) for _ in range(N)] for _ in range(N)] |
| 192 | elif mode == "Chain": |
| 193 | fixed_spatial_masks = [[1 if i == j + 1 else 0 for i in range(N)] for j in range(N)] |
no test coverage detected