AlexWortega commited on
Commit
4ac1f5f
·
verified ·
1 Parent(s): e8f907d

fix DDP hang: spawn dataloader workers (fork-after-CUDA/NCCL starved torchrun heartbeat), picklable Collate, fewer workers per rank

Browse files
Files changed (1) hide show
  1. tinyvla_b200/scripts/train_fast.py +62 -40
tinyvla_b200/scripts/train_fast.py CHANGED
@@ -96,6 +96,58 @@ def load_compatible(model, path: Path):
96
  return len(keep)
97
 
98
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99
  # ----------------------------------------------------------------------- main
100
 
101
 
@@ -205,53 +257,23 @@ def main():
205
  )
206
 
207
  # ---- pre-tokenize every distinct instruction once -----------------------
208
- from transformers import AutoTokenizer
209
-
210
- tokenizer = AutoTokenizer.from_pretrained(pcfg.lm_model_name)
211
- if src_cfg.get("type") == "hub":
212
- _tok_cache: dict = {}
213
-
214
- def _tok(tasks):
215
- new = [t for t in set(tasks) if t not in _tok_cache]
216
- if new:
217
- e = tokenizer(new, padding="max_length", truncation=True,
218
- max_length=pcfg.tokenizer_max_length, return_tensors="pt")
219
- for i, t in enumerate(new):
220
- _tok_cache[t] = (e["input_ids"][i], e["attention_mask"][i].bool())
221
- return (torch.stack([_tok_cache[t][0] for t in tasks]),
222
- torch.stack([_tok_cache[t][1] for t in tasks]))
223
- else:
224
- texts = sorted(set(source._tasks.values()))
225
- enc = tokenizer(texts, padding="max_length", truncation=True,
226
- max_length=pcfg.tokenizer_max_length, return_tensors="pt")
227
- tok_ids, tok_mask = enc["input_ids"], enc["attention_mask"].bool()
228
- tok_lookup = {t: i for i, t in enumerate(texts)}
229
- print(f"pre-tokenized {len(texts)} instructions at fixed length {pcfg.tokenizer_max_length}")
230
-
231
- def _tok(tasks):
232
- idx = torch.tensor([tok_lookup.get(t, 0) for t in tasks])
233
- return tok_ids[idx], tok_mask[idx]
234
-
235
- def collate(items):
236
- out = {}
237
- for k in items[0]:
238
- if k == "task":
239
- ids, mask = _tok([it["task"] for it in items])
240
- out["observation.language.tokens"] = ids
241
- out["observation.language.attention_mask"] = mask
242
- else:
243
- out[k] = torch.stack([it[k] for it in items])
244
- return out
245
 
 
 
 
 
 
246
  loader = torch.utils.data.DataLoader(
247
  source,
248
  batch_size=cfg["batch_size"],
249
- num_workers=cfg.get("num_workers", 12),
250
  pin_memory=True,
251
- persistent_workers=cfg.get("num_workers", 12) > 0,
252
- prefetch_factor=cfg.get("prefetch_factor", 6) if cfg.get("num_workers", 12) > 0 else None,
253
  drop_last=True,
254
  collate_fn=collate,
 
255
  )
256
 
257
  # ---- optimizer: backbone at a lower lr, exactly as train.py does --------
 
96
  return len(keep)
97
 
98
 
99
+
100
+ class Collate:
101
+ """Picklable collate: tokenizes task strings with a per-process cache.
102
+
103
+ A module-level class (not a closure) so DataLoader workers can be started
104
+ with the 'spawn' context. spawn matters under DDP: forking workers from a
105
+ process that already initialized CUDA/NCCL (with ffmpeg + rust-tokenizer
106
+ threads alive) is exactly the fork-after-CUDA hazard that hung rank
107
+ dataloaders in testing; spawned workers start clean.
108
+ """
109
+
110
+ def __init__(self, model_name: str, max_length: int):
111
+ self.model_name = model_name
112
+ self.max_length = max_length
113
+ self._tokenizer = None
114
+ self._cache: dict = {}
115
+
116
+ def _tok(self, tasks):
117
+ if self._tokenizer is None:
118
+ import os
119
+
120
+ os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
121
+ from transformers import AutoTokenizer
122
+
123
+ self._tokenizer = AutoTokenizer.from_pretrained(self.model_name)
124
+ new = [t for t in set(tasks) if t not in self._cache]
125
+ if new:
126
+ e = self._tokenizer(new, padding="max_length", truncation=True,
127
+ max_length=self.max_length, return_tensors="pt")
128
+ for i, t in enumerate(new):
129
+ self._cache[t] = (e["input_ids"][i], e["attention_mask"][i].bool())
130
+ return (torch.stack([self._cache[t][0] for t in tasks]),
131
+ torch.stack([self._cache[t][1] for t in tasks]))
132
+
133
+ def __getstate__(self):
134
+ return {"model_name": self.model_name, "max_length": self.max_length}
135
+
136
+ def __setstate__(self, st):
137
+ self.__init__(st["model_name"], st["max_length"])
138
+
139
+ def __call__(self, items):
140
+ out = {}
141
+ for k in items[0]:
142
+ if k == "task":
143
+ ids, mask = self._tok([it["task"] for it in items])
144
+ out["observation.language.tokens"] = ids
145
+ out["observation.language.attention_mask"] = mask
146
+ else:
147
+ out[k] = torch.stack([it[k] for it in items])
148
+ return out
149
+
150
+
151
  # ----------------------------------------------------------------------- main
152
 
153
 
 
257
  )
258
 
259
  # ---- pre-tokenize every distinct instruction once -----------------------
260
+ collate = Collate(pcfg.lm_model_name, pcfg.tokenizer_max_length)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
261
 
262
+ nw = cfg.get("num_workers", 12)
263
+ if ddp:
264
+ # 16 воркеров/ранг x N рангов душат CPU и heartbeat torchrun-агента
265
+ nw = cfg.get("num_workers_per_rank", max(4, nw // world))
266
+ log(f"DDP: {nw} dataloader workers per rank")
267
  loader = torch.utils.data.DataLoader(
268
  source,
269
  batch_size=cfg["batch_size"],
270
+ num_workers=nw,
271
  pin_memory=True,
272
+ persistent_workers=nw > 0,
273
+ prefetch_factor=cfg.get("prefetch_factor", 6) if nw > 0 else None,
274
  drop_last=True,
275
  collate_fn=collate,
276
+ multiprocessing_context="spawn" if nw > 0 else None,
277
  )
278
 
279
  # ---- optimizer: backbone at a lower lr, exactly as train.py does --------