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

Function get_kwargs

experiments/run_mmlu.py:136–214  ·  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

134
135
136def 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)]

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