From 87c24075038b51e7f85e981625279554553d140e Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Sun, 12 Oct 2025 20:24:43 +0800 Subject: [PATCH 1/3] more --- miles/ray/rollout_data_source.py | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/miles/ray/rollout_data_source.py b/miles/ray/rollout_data_source.py index 649246ef469..249001e4ac5 100644 --- a/miles/ray/rollout_data_source.py +++ b/miles/ray/rollout_data_source.py @@ -45,8 +45,6 @@ def __init__(self, args): self.dataset = None def get_samples(self, num_samples): - samples = [] - # TODO unify the two branches if self.dataset is not None: if self.sample_offset + num_samples <= len(self.dataset): @@ -60,14 +58,6 @@ def get_samples(self, num_samples): self.dataset.shuffle(self.epoch_id) prompt_samples += self.dataset.samples[:num_samples] self.sample_offset = num_samples - for prompt_sample in prompt_samples: - group = [] - for _ in range(self.args.n_samples_per_prompt): - sample = copy.deepcopy(prompt_sample) - sample.index = self.sample_index - self.sample_index += 1 - group.append(sample) - samples.append(group) else: for _ in range(num_samples): group = [] @@ -79,6 +69,15 @@ def get_samples(self, num_samples): group.append(sample) samples.append(group) + samples = [] + for prompt_sample in prompt_samples: + group = [] + for _ in range(self.args.n_samples_per_prompt): + sample = copy.deepcopy(prompt_sample) + sample.index = self.sample_index + self.sample_index += 1 + group.append(sample) + samples.append(group) return samples def add_samples(self, samples: list[list[Sample]]): From fb09566432ecf537278aaa32e5c3e61e1ce79f3f Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Sun, 12 Oct 2025 20:25:25 +0800 Subject: [PATCH 2/3] more --- miles/ray/rollout_data_source.py | 10 +--------- 1 file changed, 1 insertion(+), 9 deletions(-) diff --git a/miles/ray/rollout_data_source.py b/miles/ray/rollout_data_source.py index 249001e4ac5..a42cd3b23fa 100644 --- a/miles/ray/rollout_data_source.py +++ b/miles/ray/rollout_data_source.py @@ -59,15 +59,7 @@ def get_samples(self, num_samples): prompt_samples += self.dataset.samples[:num_samples] self.sample_offset = num_samples else: - for _ in range(num_samples): - group = [] - for _ in range(self.args.n_samples_per_prompt): - sample = Sample( - index=self.sample_index, - ) - self.sample_index += 1 - group.append(sample) - samples.append(group) + prompt_samples = [Sample() for _ in range(num_samples)] samples = [] for prompt_sample in prompt_samples: From 0c14664a6968608b938c2131c4bccc5332f920f5 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Sun, 12 Oct 2025 20:26:09 +0800 Subject: [PATCH 3/3] more --- miles/ray/rollout_data_source.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/miles/ray/rollout_data_source.py b/miles/ray/rollout_data_source.py index a42cd3b23fa..7c47cdd8910 100644 --- a/miles/ray/rollout_data_source.py +++ b/miles/ray/rollout_data_source.py @@ -45,7 +45,7 @@ def __init__(self, args): self.dataset = None def get_samples(self, num_samples): - # TODO unify the two branches + # TODO further improve code if self.dataset is not None: if self.sample_offset + num_samples <= len(self.dataset): prompt_samples = self.dataset.samples[self.sample_offset : self.sample_offset + num_samples]