MCPcopy Create free account
hub / github.com/FastMAS/KVCOMM / get_kwargs

Function get_kwargs

experiments/benchmark_TTFT.py:180–258  ·  view source on GitHub ↗
(
    mode: Union[
        Literal["DirectAnswer"],
        Literal["FullConnected"],
        Literal["Random"],
        Literal["Chain"],
        Literal["Debate"],
        Literal["Layered"],
        Literal["Star"],
        Literal["Mesh"],
    ],
    N: int,
)

Source from the content-addressed store, hash-verified

178
179
180def 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)]

Callers 1

mainFunction · 0.70

Calls 3

generate_layered_graphFunction · 0.70
generate_mesh_graphFunction · 0.70
generate_star_graphFunction · 0.70

Tested by

no test coverage detected