| 964 | |
| 965 | template <typename T> |
| 966 | void DatasetImpl<T>::DynamicAdjustChannelNum(int channel_num, |
| 967 | bool discard_remaining_ins) { |
| 968 | if (channel_num_ == channel_num) { |
| 969 | VLOG(3) << "DatasetImpl<T>::DynamicAdjustChannelNum channel_num_=" |
| 970 | << channel_num_ << ", channel_num_=channel_num, no need to adjust"; |
| 971 | return; |
| 972 | } |
| 973 | VLOG(3) << "adjust channel num from " << channel_num_ << " to " |
| 974 | << channel_num; |
| 975 | channel_num_ = channel_num; |
| 976 | std::vector<paddle::framework::Channel<T>>* origin_channels = nullptr; |
| 977 | std::vector<paddle::framework::Channel<T>>* other_channels = nullptr; |
| 978 | std::vector<paddle::framework::Channel<PvInstance>>* origin_pv_channels = |
| 979 | nullptr; |
| 980 | std::vector<paddle::framework::Channel<PvInstance>>* other_pv_channels = |
| 981 | nullptr; |
| 982 | |
| 983 | // find out which channel (output or consume) has data |
| 984 | int cur_channel = 0; |
| 985 | uint64_t output_channels_data_size = 0; |
| 986 | uint64_t consume_channels_data_size = 0; |
| 987 | PADDLE_ENFORCE_EQ(multi_output_channel_.size(), |
| 988 | multi_consume_channel_.size(), |
| 989 | common::errors::InvalidArgument( |
| 990 | "The size of multi_output_channel (%d) does not match " |
| 991 | "the size of multi_consume_channel (%d).", |
| 992 | multi_output_channel_.size(), |
| 993 | multi_consume_channel_.size())); |
| 994 | |
| 995 | for (size_t i = 0; i < multi_output_channel_.size(); ++i) { |
| 996 | output_channels_data_size += multi_output_channel_[i]->Size(); |
| 997 | consume_channels_data_size += multi_consume_channel_[i]->Size(); |
| 998 | } |
| 999 | |
| 1000 | if (output_channels_data_size != 0) { |
| 1001 | PADDLE_ENFORCE_EQ(consume_channels_data_size, |
| 1002 | 0, |
| 1003 | common::errors::InvalidArgument( |
| 1004 | "When output_channels_data_size (%d) is not zero, " |
| 1005 | "consume_channels_data_size (%d) should be zero.", |
| 1006 | output_channels_data_size, |
| 1007 | consume_channels_data_size)); |
| 1008 | cur_channel = 0; |
| 1009 | } else { |
| 1010 | PADDLE_ENFORCE_EQ( |
| 1011 | output_channels_data_size, |
| 1012 | 0, |
| 1013 | common::errors::InvalidArgument( |
| 1014 | "When output_channels_data_size is zero, it should be zero. " |
| 1015 | "consume_channels_data_size: %d", |
| 1016 | consume_channels_data_size)); |
| 1017 | cur_channel = 1; |
| 1018 | } |
| 1019 | |
| 1020 | if (cur_channel == 0) { // NOLINT |
| 1021 | origin_channels = &multi_output_channel_; |
| 1022 | other_channels = &multi_consume_channel_; |
| 1023 | origin_pv_channels = &multi_pv_output_; |