| 100 | } |
| 101 | |
| 102 | TGpuSamplesGrouping<NCudaLib::TStripeMapping> TGpuSamplesGroupingHelper<NCudaLib::TMirrorMapping>::MakeStripeGrouping(const TGpuSamplesGrouping<NCudaLib::TMirrorMapping>& mirrorMapping, |
| 103 | const TCudaBuffer<const ui32, NCudaLib::TStripeMapping>& indices) { |
| 104 | CB_ENSURE(indices.GetObjectsSlice() == mirrorMapping.CurrentDocsSlice); |
| 105 | TSlice docsSlice = indices.GetObjectsSlice(); |
| 106 | docsSlice.Left += mirrorMapping.CurrentDocsSlice.Left; |
| 107 | docsSlice.Right += mirrorMapping.CurrentDocsSlice.Left; |
| 108 | const IQueriesGrouping& grouping = *mirrorMapping.Grouping; |
| 109 | const ui32 queryIdOffset = grouping.GetQueryId(docsSlice.Left); |
| 110 | |
| 111 | const NCudaLib::TDistributedObject<ui32>& baseBias = mirrorMapping.GetOffsetsBias(); |
| 112 | NCudaLib::TDistributedObject<ui32> deviceOffsetsBiases = NCudaLib::GetCudaManager().CreateDistributedObject<ui32>(0); |
| 113 | |
| 114 | TVector<TSlice> groupMetaSlices(deviceOffsetsBiases.DeviceCount()); |
| 115 | |
| 116 | ui32 offset = 0; |
| 117 | for (ui32 dev = 0; dev < deviceOffsetsBiases.DeviceCount(); ++dev) { |
| 118 | TSlice deviceSlice = indices.GetMapping().DeviceSlice(dev); |
| 119 | const ui32 localBias = baseBias.At(dev) + offset; |
| 120 | deviceOffsetsBiases.Set(dev, localBias); |
| 121 | ui32 firstGroupId = grouping.GetQueryId(docsSlice.Left + offset) - queryIdOffset; |
| 122 | ui32 lastGroupId = grouping.GetQueryId(docsSlice.Left + offset + deviceSlice.Size()) - queryIdOffset; |
| 123 | offset += deviceSlice.Size(); |
| 124 | groupMetaSlices[dev] = TSlice(firstGroupId, lastGroupId); |
| 125 | } |
| 126 | |
| 127 | NCudaLib::TStripeMapping groupMetaMapping(std::move(groupMetaSlices)); |
| 128 | CB_ENSURE(groupMetaMapping.GetObjectsSlice() == TSlice(0, grouping.GetQueryId(docsSlice.Right) - queryIdOffset)); |
| 129 | NCudaLib::TCudaBuffer<const ui32, NCudaLib::TStripeMapping> offsets = NCudaLib::StripeView(mirrorMapping.Offsets, groupMetaMapping); |
| 130 | NCudaLib::TCudaBuffer<const ui32, NCudaLib::TStripeMapping> sizes = NCudaLib::StripeView(mirrorMapping.Sizes, groupMetaMapping); |
| 131 | |
| 132 | TGpuSamplesGrouping<NCudaLib::TStripeMapping> samplesGrouping(&grouping, |
| 133 | mirrorMapping.CurrentDocsSlice, |
| 134 | std::move(offsets), |
| 135 | std::move(sizes), |
| 136 | std::move(deviceOffsetsBiases)); |
| 137 | |
| 138 | { |
| 139 | const TQueriesGrouping* pointwiseSamplesGrouping = dynamic_cast<const TQueriesGrouping*>(&grouping); |
| 140 | if (pointwiseSamplesGrouping && pointwiseSamplesGrouping->GetFlatQueryPairs().size()) { |
| 141 | NCudaLib::TStripeMapping pairsMapping = groupMetaMapping.Transform([&](const TSlice& groupsSlice) -> ui64 { |
| 142 | ui32 firstGroupId = groupsSlice.Left + queryIdOffset; |
| 143 | ui32 lastGroupId = groupsSlice.Right + queryIdOffset; |
| 144 | return pointwiseSamplesGrouping->GetQueryPairOffset(lastGroupId) - pointwiseSamplesGrouping->GetQueryPairOffset(firstGroupId); |
| 145 | }); |
| 146 | samplesGrouping.Pairs = NCudaLib::StripeView(mirrorMapping.Pairs, pairsMapping); |
| 147 | samplesGrouping.PairsWeights = NCudaLib::StripeView(mirrorMapping.PairsWeights, pairsMapping); |
| 148 | } |
| 149 | } |
| 150 | return samplesGrouping; |
| 151 | } |
| 152 | |
| 153 | TGpuSamplesGrouping<NCudaLib::TStripeMapping> |
| 154 | TGpuSamplesGroupingHelper<NCudaLib::TStripeMapping>::CreateGpuGrouping(const TDocParallelDataSet& dataSet) { |
nothing calls this directly
no test coverage detected