| 7 | |
| 8 | namespace NCatboostCuda { |
| 9 | TFeatureParallelDataSetsHolder TFeatureParallelDataSetHoldersBuilder::BuildDataSet(const ui32 permutationCount, |
| 10 | NPar::ILocalExecutor* localExecutor) { |
| 11 | TFeatureParallelDataSetsHolder dataSetsHolder(DataProvider, |
| 12 | FeaturesManager); |
| 13 | |
| 14 | Y_ASSERT(dataSetsHolder.CompressedIndex); |
| 15 | TSharedCompressedIndexBuilder<TDataSetLayout> compressedIndexBuilder(*dataSetsHolder.CompressedIndex, |
| 16 | localExecutor); |
| 17 | |
| 18 | dataSetsHolder.CtrTargets = BuildCtrTarget(FeaturesManager, |
| 19 | DataProvider, |
| 20 | LinkedTest); |
| 21 | auto& ctrsTarget = *dataSetsHolder.CtrTargets; |
| 22 | |
| 23 | { |
| 24 | dataSetsHolder.LearnCatFeaturesDataSet = MakeHolder<TCompressedCatFeatureDataSet>(CatFeaturesStorage); |
| 25 | BuildCompressedCatFeatures(DataProvider, |
| 26 | *dataSetsHolder.LearnCatFeaturesDataSet, |
| 27 | localExecutor); |
| 28 | |
| 29 | if (LinkedTest) { |
| 30 | dataSetsHolder.TestCatFeaturesDataSet = MakeHolder<TCompressedCatFeatureDataSet>(CatFeaturesStorage); |
| 31 | BuildCompressedCatFeatures(*LinkedTest, |
| 32 | *dataSetsHolder.TestCatFeaturesDataSet, |
| 33 | localExecutor); |
| 34 | } |
| 35 | } |
| 36 | |
| 37 | TAtomicSharedPtr<TPermutationScope> permutationIndependentScope = new TPermutationScope; |
| 38 | |
| 39 | dataSetsHolder.PermutationDataSets.resize(permutationCount); |
| 40 | |
| 41 | const auto learnWeights = NCB::GetWeights(*DataProvider.TargetData); |
| 42 | |
| 43 | const bool isTrivialLearnWeights = AreEqualTo(learnWeights, 1.0f); |
| 44 | { |
| 45 | const auto learnMapping = NCudaLib::TMirrorMapping(ctrsTarget.LearnSlice.Size()); |
| 46 | |
| 47 | if (isTrivialLearnWeights == ctrsTarget.IsTrivialWeights()) { |
| 48 | dataSetsHolder.DirectWeights = ctrsTarget.Weights.SliceView(ctrsTarget.LearnSlice); |
| 49 | } else { |
| 50 | dataSetsHolder.DirectWeights.Reset(learnMapping); |
| 51 | dataSetsHolder.DirectWeights.Write(learnWeights); |
| 52 | } |
| 53 | if (isTrivialLearnWeights && ctrsTarget.IsTrivialWeights()) { |
| 54 | dataSetsHolder.DirectTarget = ctrsTarget.WeightedTarget.SliceView(ctrsTarget.LearnSlice); |
| 55 | } else { |
| 56 | dataSetsHolder.DirectTarget.Reset(learnMapping); |
| 57 | dataSetsHolder.DirectTarget.Write(*DataProvider.TargetData->GetOneDimensionalTarget()); |
| 58 | } |
| 59 | } |
| 60 | |
| 61 | for (ui32 permutationId = 0; permutationId < permutationCount; ++permutationId) { |
| 62 | TDataPermutation permutation = NCatboostCuda::GetPermutation(DataProvider, |
| 63 | permutationId, |
| 64 | DataProviderPermutationBlockSize); |
| 65 | |
| 66 | const auto targetsMapping = NCudaLib::TMirrorMapping(ctrsTarget.LearnSlice.Size()); |
nothing calls this directly
no test coverage detected