fix: rewrite GRPO data pipeline for offline record-level access

- process_list_field returns List[List[int]] preserving per-response boundaries
- GRPODataset rewritten to record-level __getitem__ (no windowing/stride)
- grpo_collate_fn pads variable-length responses into [B, G, R] tensors
- JsonlStore detects nested List[List[int]] and stores List[Tensor] per record
- Store._normalize skips nested-list keys from cumsum bookkeeping
- Pipeline._flush handles nested lists without cross-record flattening
- Export grpo_collate_fn from astrai.dataset
- 6 new GRPO tests + 2 updated builder tests, 114 total pass
This commit is contained in:
2026-07-17 14:34:41 +08:00
parent c17aa0dc54
commit a1ea26d367
7 changed files with 502 additions and 75 deletions
+16 -3
View File
@@ -349,7 +349,17 @@ def test_grpo_basic(chat_tokenizer, builder):
assert "responses" in result
assert "masks" in result
assert "rewards" in result
assert len(result["responses"]) == len(result["masks"])
# responses is List[List[int]] — one per response
assert len(result["responses"]) == 4
assert all(isinstance(r, list) for r in result["responses"])
assert all(isinstance(r[0], int) for r in result["responses"])
# masks is List[List[int]] — one per response, matching length
assert len(result["masks"]) == 4
for i in range(4):
assert len(result["masks"][i]) == len(result["responses"][i])
assert result["rewards"] == [1.0, 0.5, 0.8, 0.2]
@@ -362,8 +372,11 @@ def test_grpo_response_tokens_all_trained(chat_tokenizer, builder):
}
result = builder.build(item, config, chat_tokenizer)
masks = result["masks"]
assert all(m == 1 for m in masks)
assert len(masks) == len(result["responses"])
# masks is List[List[int]] — each response's mask should be all 1s
assert len(masks) == 2
for m in masks:
assert all(v == 1 for v in m)
assert len(m) == len(result["responses"][masks.index(m)])
def test_grpo_single_reward(chat_tokenizer, builder):