| 75 | |
| 76 | |
| 77 | def req_inference( |
| 78 | endpoint: str, |
| 79 | inference_parallel_size: int, |
| 80 | uds: str | None = None, |
| 81 | ) -> Callable[[list[tuple[str, str]]], None]: |
| 82 | rank = int(os.getenv("RANK", None)) |
| 83 | src = rank // inference_parallel_size * inference_parallel_size |
| 84 | |
| 85 | def req_func(socket_paths: list[tuple[str, str]]): |
| 86 | if rank == src: |
| 87 | request_inference_to_update( |
| 88 | f"{endpoint}/collective_rpc", |
| 89 | dict(socket_paths[src : src + inference_parallel_size]), |
| 90 | uds=uds, |
| 91 | ) |
| 92 | |
| 93 | return req_func |
| 94 | |
| 95 | |
| 96 | def update_weights( |