# Copyright (c) 2025, HUAWEI CORPORATION.  All rights reserved.

import ray
import numpy as np

from mindspeed_rl.utils.loggers import Loggers
from mindspeed_rl.utils.utils import get_current_dp_range_indexes, extract_from_dict, is_multimodal
from mindspeed_rl.utils.pad_process import remove_padding_tensor_dict_to_dict, padding_dict_to_tensor_dict

logger = Loggers("DynamicSampling")


@ray.remote
class DynamicSampling(object):
    """
    DynamicSampling class for filtering high-quality training data.
    
    This class implements dynamic sampling logic that filters experience data
    based on reward model score variance within each prompt group, keeping only
    groups with diverse enough rewards for effective training.
    """

    def initialize(self, megatron_config, rl_config):
        """
        Initialize the DynamicSampling worker.
        
        Args:
            megatron_config: Configuration for Megatron-LM.
            rl_config: Configuration for reinforcement learning.
        """
        # Reinforcement learning configuration
        self.rl_config = rl_config
        # Megatron-LM configuration
        self.megatron_config = megatron_config
        # Number of samples generated per prompt
        self.n_samples_per_prompt = rl_config.n_samples_per_prompt
        # Global batch size for training
        self.global_batch_size = megatron_config.global_batch_size
        # Whether to guarantee data order during sampling
        self.guarantee_order = rl_config.guarantee_order

    def init_transfer_dock(self, td, mm_td=None, sampling_transfer_dock=None, mm_sampling_transfer_dock=None):
        """
        Initialize transfer dock references for data communication.
        """
        self.td = td
        self.mm_td = mm_td
        self.sampling_transfer_dock = sampling_transfer_dock
        self.mm_sampling_transfer_dock = mm_sampling_transfer_dock

    def dynamic_sampling(self):
        """
        Perform dynamic sampling to filter high-quality training data.
        """
        experience_consumer_stage = 'dynamic_sampling'
        experience_columns = ['prompts', 'prompt_length', 'responses', 'response_length', 'input_ids', 'rm_scores',
                              'metric_for_dapo', *self.megatron_config.dataset_additional_keys]
        if self.rl_config.multi_turn_enable:
            experience_columns.extend(['response_mask', 'tool_call_num'])
        experience_count = self.rl_config.dynamic_sampling_dispatch_size
        assign_batch_size = self.global_batch_size * self.n_samples_per_prompt
        sorted_indexes = get_current_dp_range_indexes(experience_count=experience_count,
                                                      assign_batch_size=assign_batch_size) if self.guarantee_order else None
        sampling_td = self.sampling_transfer_dock
        main_td = self.td
        mm_sampling_td = self.mm_sampling_transfer_dock
        mm_main_td = self.mm_td

        while not ray.get(sampling_td.all_consumed.remote(experience_consumer_stage)):
            batch_data, index = ray.get(
                sampling_td.get_experience.remote(
                    experience_consumer_stage,
                    experience_columns,
                    experience_count,
                    indexes=sorted_indexes.pop(0) if self.guarantee_order else None,
                    get_n_samples=True
                )
            )
            batch_data = remove_padding_tensor_dict_to_dict(batch_data)
            if batch_data and index:
                metric_values = batch_data['metric_for_dapo']
                kept_idx_list = []
                for idx in range(0, experience_count, self.n_samples_per_prompt):
                    metric_group = metric_values[idx: idx + self.n_samples_per_prompt]
                    if np.std(metric_group) > 0 or len(metric_group) == 1:
                        kept_idx_list.extend(list(range(idx, idx + self.n_samples_per_prompt)))
                if not kept_idx_list:
                    logger.info("dynamic_sampling: kept_idx_list is empty")
                    continue

                index_list = ray.get(main_td.prefetch_request_index.remote(len(kept_idx_list)))
                if index_list:
                    kept_idx_list = kept_idx_list[:len(index_list)]
                    experience_data = extract_from_dict(batch_data, kept_idx_list)
                    experience_data = padding_dict_to_tensor_dict(experience_data)
                    ray.get(main_td.put_experience.remote(experience_data, index_list))
                    if is_multimodal() and mm_sampling_td and mm_main_td:
                        mm_columns = ray.get(mm_sampling_td.get_columns.remote(experience_consumer_stage))
                        batch_mm_data = ray.get(
                            mm_sampling_td.get_experience_dict.remote(mm_columns, kept_idx_list, False)
                        )
                        mm_index_list = [i // self.n_samples_per_prompt for i in index_list]
                        ray.get(mm_main_td.put_experience.remote(batch_mm_data, mm_index_list))