| 47 | } |
| 48 | |
| 49 | TGpuSamplesGrouping<NCudaLib::TMirrorMapping> |
| 50 | TGpuSamplesGroupingHelper<NCudaLib::TMirrorMapping>::SliceGrouping(const TGpuSamplesGrouping<NCudaLib::TMirrorMapping>& grouping, |
| 51 | const TSlice& localSlice) { |
| 52 | CB_ENSURE(localSlice.Size() <= grouping.CurrentDocsSlice.Size()); |
| 53 | TSlice globalSlice; |
| 54 | globalSlice.Left = localSlice.Left + grouping.CurrentDocsSlice.Left; |
| 55 | globalSlice.Right = localSlice.Right + grouping.CurrentDocsSlice.Left; |
| 56 | |
| 57 | const ui32 firstGroupId = grouping.Grouping->GetQueryId(globalSlice.Left); |
| 58 | const ui32 lastGroupId = grouping.Grouping->GetQueryId(globalSlice.Right); |
| 59 | |
| 60 | const ui32 firstGroupDoc = grouping.Grouping->GetQueryOffset(firstGroupId); |
| 61 | const ui32 lastGroupDoc = grouping.Grouping->GetQueryOffset(lastGroupId); |
| 62 | |
| 63 | CB_ENSURE(firstGroupDoc == globalSlice.Left, "Error: slice should be group-consistent"); |
| 64 | CB_ENSURE(lastGroupDoc == globalSlice.Right, "Error: slice should be group-consistent"); |
| 65 | |
| 66 | auto biases = NCudaLib::GetCudaManager().CreateDistributedObject<ui32>(0); |
| 67 | for (ui32 dev = 0; dev < biases.DeviceCount(); ++dev) { |
| 68 | biases.Set(dev, firstGroupDoc); |
| 69 | } |
| 70 | |
| 71 | TSlice groupsSlice; |
| 72 | const ui32 groupOffset = grouping.Grouping->GetQueryId(grouping.CurrentDocsSlice.Left); |
| 73 | CB_ENSURE(firstGroupId >= groupOffset); |
| 74 | groupsSlice.Left = firstGroupId - groupOffset; |
| 75 | groupsSlice.Right = lastGroupId - groupOffset; |
| 76 | CB_ENSURE(grouping.Offsets.GetObjectsSlice() == TSlice(0, grouping.Grouping->GetQueryId( |
| 77 | grouping.CurrentDocsSlice.Right) - |
| 78 | groupOffset)); |
| 79 | |
| 80 | TGpuSamplesGrouping<NCudaLib::TMirrorMapping> sliceGrouping(grouping.Grouping, |
| 81 | globalSlice, |
| 82 | grouping.Offsets.SliceView(groupsSlice), |
| 83 | grouping.Sizes.SliceView(groupsSlice), |
| 84 | std::move(biases)); |
| 85 | { |
| 86 | const TQueriesGrouping* pointwiseSamplesGrouping = dynamic_cast<const TQueriesGrouping*>(grouping.Grouping); |
| 87 | if (pointwiseSamplesGrouping && pointwiseSamplesGrouping->GetFlatQueryPairs().size()) { |
| 88 | TSlice pairsSlice; |
| 89 | CB_ENSURE(firstGroupId >= groupOffset); |
| 90 | const ui32 shift = pointwiseSamplesGrouping->GetQueryPairOffset(groupOffset); |
| 91 | pairsSlice.Left = pointwiseSamplesGrouping->GetQueryPairOffset(firstGroupId) - shift; |
| 92 | pairsSlice.Right = pointwiseSamplesGrouping->GetQueryPairOffset(lastGroupId) - shift; |
| 93 | |
| 94 | sliceGrouping.Pairs = grouping.Pairs.SliceView(pairsSlice); |
| 95 | sliceGrouping.PairsWeights = grouping.PairsWeights.SliceView(pairsSlice); |
| 96 | } |
| 97 | } |
| 98 | |
| 99 | return sliceGrouping; |
| 100 | } |
| 101 | |
| 102 | TGpuSamplesGrouping<NCudaLib::TStripeMapping> TGpuSamplesGroupingHelper<NCudaLib::TMirrorMapping>::MakeStripeGrouping(const TGpuSamplesGrouping<NCudaLib::TMirrorMapping>& mirrorMapping, |
| 103 | const TCudaBuffer<const ui32, NCudaLib::TStripeMapping>& indices) { |
nothing calls this directly
no test coverage detected