Aditya Guntur commited on
Commit
5d330b7
·
1 Parent(s): 29d6757

fix(rollout): replace generate_rollout_completions (vLLM-only) with direct model.generate()

Browse files

generate_rollout_completions hard-crashes when use_vllm=False. Replace with
_generate_no_vllm() which uses HF model.generate() + output_scores=True to
get prompt_ids, completion_ids, logprobs, and text without vLLM.

Files changed (1) hide show
  1. training/rollout.py +51 -5
training/rollout.py CHANGED
@@ -16,11 +16,59 @@ import json
16
  import re
17
  from typing import Any
18
 
19
- from trl.experimental.openenv import generate_rollout_completions
 
20
 
21
  from training.dataset import parse_seed_from_prompt
22
  from training.prompts import SYSTEM_PROMPT, format_observation
23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  MAX_STEPS = 40
25
  # ~3000 tokens at 4 chars/token; leaves room for completion tokens
26
  MAX_PROMPT_CHARS = 12_000
@@ -209,14 +257,12 @@ def rollout_once(
209
  enable_thinking=False,
210
  )
211
 
212
- rollout_out = generate_rollout_completions(trainer, [prompt_text])[0]
213
  prompt_ids.extend(rollout_out["prompt_ids"])
214
  completion_ids.extend(rollout_out["completion_ids"])
215
  logprobs.extend(rollout_out["logprobs"])
216
 
217
- completion_text = rollout_out.get("text") or tokenizer.decode(
218
- rollout_out["completion_ids"], skip_special_tokens=True
219
- )
220
 
221
  # Parse action; fall back gracefully on parse failure
222
  parsed = extract_json_action(completion_text)
 
16
  import re
17
  from typing import Any
18
 
19
+ import torch
20
+ import torch.nn.functional as F
21
 
22
  from training.dataset import parse_seed_from_prompt
23
  from training.prompts import SYSTEM_PROMPT, format_observation
24
 
25
+
26
+ # ---------------------------------------------------------------------------
27
+ # HF model.generate() — replaces generate_rollout_completions (vLLM-only)
28
+ # ---------------------------------------------------------------------------
29
+
30
+ def _generate_no_vllm(trainer, prompt_text: str, tokenizer, max_new_tokens: int = 512) -> dict:
31
+ """Generate one completion using HF model.generate() without vLLM.
32
+
33
+ Returns the same dict shape as generate_rollout_completions so the rest
34
+ of rollout_once is unchanged:
35
+ prompt_ids: list[int]
36
+ completion_ids: list[int]
37
+ logprobs: list[float] (per-token log-prob under current policy)
38
+ text: str
39
+ """
40
+ device = next(trainer.model.parameters()).device
41
+ enc = tokenizer(prompt_text, return_tensors="pt").to(device)
42
+ prompt_len = enc["input_ids"].shape[1]
43
+
44
+ with torch.no_grad():
45
+ out = trainer.model.generate(
46
+ **enc,
47
+ max_new_tokens=max_new_tokens,
48
+ do_sample=True,
49
+ temperature=0.7,
50
+ pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
51
+ output_scores=True,
52
+ return_dict_in_generate=True,
53
+ )
54
+
55
+ prompt_ids = enc["input_ids"][0].tolist()
56
+ completion_ids = out.sequences[0][prompt_len:].tolist()
57
+
58
+ # Per-token log-probs from output.scores (one score tensor per new token)
59
+ logprobs = [
60
+ F.log_softmax(score[0], dim=-1)[tok_id].item()
61
+ for score, tok_id in zip(out.scores, completion_ids)
62
+ ]
63
+
64
+ text = tokenizer.decode(completion_ids, skip_special_tokens=True)
65
+ return {
66
+ "prompt_ids": prompt_ids,
67
+ "completion_ids": completion_ids,
68
+ "logprobs": logprobs,
69
+ "text": text,
70
+ }
71
+
72
  MAX_STEPS = 40
73
  # ~3000 tokens at 4 chars/token; leaves room for completion tokens
74
  MAX_PROMPT_CHARS = 12_000
 
257
  enable_thinking=False,
258
  )
259
 
260
+ rollout_out = _generate_no_vllm(trainer, prompt_text, tokenizer)
261
  prompt_ids.extend(rollout_out["prompt_ids"])
262
  completion_ids.extend(rollout_out["completion_ids"])
263
  logprobs.extend(rollout_out["logprobs"])
264
 
265
+ completion_text = rollout_out["text"]
 
 
266
 
267
  # Parse action; fall back gracefully on parse failure
268
  parsed = extract_json_action(completion_text)