MCPcopy Create free account
hub / github.com/catboost/catboost / SliceGrouping

Method SliceGrouping

catboost/cuda/gpu_data/samples_grouping_gpu.cpp:49–100  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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) {

Callers

nothing calls this directly

Calls 11

DeviceCountMethod · 0.80
SliceViewMethod · 0.80
TSliceClass · 0.50
moveFunction · 0.50
SizeMethod · 0.45
GetQueryIdMethod · 0.45
GetQueryOffsetMethod · 0.45
SetMethod · 0.45
GetObjectsSliceMethod · 0.45
sizeMethod · 0.45
GetQueryPairOffsetMethod · 0.45

Tested by

no test coverage detected