(self,
*from_pretrained_args,
gpus:Optional[int]=None, gpu_ids:Optional[List[int]]=None,
**from_pretrained_kwargs)
| 170 | >>> wrapper.generate(["The best programming language is"], max_gen_len=10, temperature=0) |
| 171 | """ |
| 172 | def __init__(self, |
| 173 | *from_pretrained_args, |
| 174 | gpus:Optional[int]=None, gpu_ids:Optional[List[int]]=None, |
| 175 | **from_pretrained_kwargs): |
| 176 | if gpus is None and gpu_ids is None: |
| 177 | raise ValueError('You must specify either gpus or gpu_ids') |
| 178 | |
| 179 | if gpu_ids is None: |
| 180 | gpu_ids = list(range(gpus)) |
| 181 | |
| 182 | self.outer_world = _save_world() |
| 183 | _reset_world() |
| 184 | |
| 185 | # launch model worker subprocesses and create inner world |
| 186 | port = accessory.util.misc.find_free_port(10000+int(os.getpid())%100*100, 10100+int(os.getpid())%100*100) |
| 187 | print(f"Launching {len(gpu_ids)} processes for hosting model with model parallel size {len(gpu_ids)}") |
| 188 | print("Note that only the output from the FIRST model process will be printed") |
| 189 | self.model_workers = [] |
| 190 | for rank in range(1, len(gpu_ids) + 1): |
| 191 | p = subprocess.Popen([sys.executable, __file__, str(port), str(rank), str(len(gpu_ids)+1)]) |
| 192 | self.model_workers.append(p) |
| 193 | |
| 194 | atexit.register(self.on_exit) |
| 195 | |
| 196 | store = dist.TCPStore("127.0.0.1", port, len(gpu_ids) + 1, True) |
| 197 | dist.init_process_group( |
| 198 | backend="gloo", rank=0, world_size=len(gpu_ids) + 1, |
| 199 | # init_method=f"tcp://127.0.0.1:{port}", |
| 200 | store=store |
| 201 | ) |
| 202 | self.inner_world = _save_world() |
| 203 | |
| 204 | |
| 205 | dist.broadcast_object_list([[gpu_ids, from_pretrained_args, from_pretrained_kwargs]], src=0) |
| 206 | dist.new_group(ranks=list(range(1, len(gpu_ids) + 1)), backend="nccl") |
| 207 | dist.barrier() |
| 208 | |
| 209 | _load_world(*self.outer_world) |
| 210 | |
| 211 | def compute_logits(self, *args, **kwargs): |
| 212 | """ |
nothing calls this directly
no test coverage detected