diff --git a/miles/rollout/generate_utils/openai_endpoint_utils.py b/miles/rollout/generate_utils/openai_endpoint_utils.py index 645742afb37..f1cb76f1565 100644 --- a/miles/rollout/generate_utils/openai_endpoint_utils.py +++ b/miles/rollout/generate_utils/openai_endpoint_utils.py @@ -7,7 +7,7 @@ import logging import random from argparse import Namespace -from copy import deepcopy +from copy import copy, deepcopy from typing import Any import httpx @@ -743,7 +743,10 @@ def _compute_sample_from_openai_record( output_token_ids = [item[1] for item in choice["meta_info"]["output_token_logprobs"]] output_log_probs = [item[0] for item in choice["meta_info"]["output_token_logprobs"]] - sample = deepcopy(input_sample) + sample = copy(input_sample) + sample.metadata = dict(input_sample.metadata) + sample.weight_versions = list(input_sample.weight_versions) + sample.prefix_cache_info = deepcopy(input_sample.prefix_cache_info) request_input_ids = record.request.get("input_ids") if request_input_ids is not None: assert ( diff --git a/miles/rollout/generate_utils/sample_utils.py b/miles/rollout/generate_utils/sample_utils.py index b909551e564..5844f09b91b 100644 --- a/miles/rollout/generate_utils/sample_utils.py +++ b/miles/rollout/generate_utils/sample_utils.py @@ -1,4 +1,4 @@ -from copy import deepcopy +from copy import copy from dataclasses import fields from miles.utils.types import Sample @@ -30,7 +30,7 @@ def merge_samples(samples: list[Sample], tokenizer) -> Sample: def _merge_sample_pair(a: Sample, b: Sample, tokenizer) -> Sample: """Merge two samples generated from sibling inference engine calls.""" - a, b = deepcopy(a), deepcopy(b) + a, b = copy(a), copy(b) def _merge_equal_value(field): x = getattr(a, field)