Skip to content

Commit bd7f342

Browse files
committed
make it pass tests
1 parent 9740d9a commit bd7f342

4 files changed

Lines changed: 59 additions & 70 deletions

File tree

Framework/Core/include/Framework/ASoA.h

Lines changed: 32 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -502,13 +502,13 @@ class ColumnIterator : ChunkingPolicy
502502
: mColumn{column},
503503
mCurrent{nullptr},
504504
mCurrentPos{nullptr},
505+
mGlobalOffset{nullptr},
505506
mLast{nullptr},
506507
mFirstIndex{0},
507-
mCurrentChunk{0},
508-
mOffset{0}
508+
mCurrentChunk{0}
509509
{
510510
auto array = getCurrentArray();
511-
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) + (mOffset >> SCALE_FACTOR);
511+
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data());
512512
mLast = mCurrent + array->length();
513513
}
514514

@@ -524,21 +524,19 @@ class ColumnIterator : ChunkingPolicy
524524
{
525525
auto previousArray = getCurrentArray();
526526
mFirstIndex += previousArray->length();
527-
528527
mCurrentChunk++;
529528
auto array = getCurrentArray();
530-
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) + (mOffset >> SCALE_FACTOR) - (mFirstIndex >> SCALE_FACTOR);
529+
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) - (mFirstIndex >> SCALE_FACTOR);
531530
mLast = mCurrent + array->length() + (mFirstIndex >> SCALE_FACTOR);
532531
}
533532

534533
void prevChunk() const
535534
{
536535
auto previousArray = getCurrentArray();
537536
mFirstIndex -= previousArray->length();
538-
539537
mCurrentChunk--;
540538
auto array = getCurrentArray();
541-
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) + (mOffset >> SCALE_FACTOR) - (mFirstIndex >> SCALE_FACTOR);
539+
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) - (mFirstIndex >> SCALE_FACTOR);
542540
mLast = mCurrent + array->length() + (mFirstIndex >> SCALE_FACTOR);
543541
}
544542

@@ -561,24 +559,24 @@ class ColumnIterator : ChunkingPolicy
561559
mCurrentChunk = mColumn->num_chunks() - 1;
562560
auto array = getCurrentArray();
563561
mFirstIndex = mColumn->length() - array->length();
564-
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) + (mOffset >> SCALE_FACTOR) - (mFirstIndex >> SCALE_FACTOR);
562+
mCurrent = reinterpret_cast<unwrap_t<T> const*>(array->values()->data()) - (mFirstIndex >> SCALE_FACTOR);
565563
mLast = mCurrent + array->length() + (mFirstIndex >> SCALE_FACTOR);
566564
}
567565

568566
auto operator*() const
569567
requires std::same_as<bool, std::decay_t<T>>
570568
{
571569
checkSkipChunk();
572-
return (*(mCurrent - (mOffset >> SCALE_FACTOR) + ((*mCurrentPos + mOffset) >> SCALE_FACTOR)) & (1 << ((*mCurrentPos + mOffset) & 0x7))) != 0;
570+
return (*(mCurrent + ((*mCurrentPos + *mGlobalOffset) >> SCALE_FACTOR)) & (1 << ((*mCurrentPos + *mGlobalOffset) & ((1 << SCALE_FACTOR) - 1)))) != 0;
573571
}
574572

575573
auto operator*() const
576574
requires((!std::same_as<bool, std::decay_t<T>>) && std::same_as<arrow_array_for_t<T>, arrow::ListArray>)
577575
{
578576
checkSkipChunk();
579577
auto list = std::static_pointer_cast<arrow::ListArray>(mColumn->chunk(mCurrentChunk));
580-
auto offset = list->value_offset(*mCurrentPos - mFirstIndex);
581-
auto length = list->value_length(*mCurrentPos - mFirstIndex);
578+
auto offset = list->value_offset(*mCurrentPos + *mGlobalOffset - mFirstIndex);
579+
auto length = list->value_length(*mCurrentPos + *mGlobalOffset - mFirstIndex);
582580
return gsl::span<unwrap_t<T> const>{mCurrent + mFirstIndex + offset, mCurrent + mFirstIndex + (offset + length)};
583581
}
584582

@@ -587,14 +585,14 @@ class ColumnIterator : ChunkingPolicy
587585
{
588586
checkSkipChunk();
589587
auto array = std::static_pointer_cast<arrow::BinaryViewArray>(mColumn->chunk(mCurrentChunk));
590-
return array->GetView(*mCurrentPos - mFirstIndex);
588+
return array->GetView(*mCurrentPos + *mGlobalOffset - mFirstIndex);
591589
}
592590

593591
decltype(auto) operator*() const
594592
requires((!std::same_as<bool, std::decay_t<T>>) && !std::same_as<arrow_array_for_t<T>, arrow::ListArray> && !std::same_as<arrow_array_for_t<T>, arrow::BinaryViewArray>)
595593
{
596594
checkSkipChunk();
597-
return *(mCurrent + (*mCurrentPos >> SCALE_FACTOR));
595+
return *(mCurrent + ((*mCurrentPos + *mGlobalOffset) >> SCALE_FACTOR));
598596
}
599597

600598
// Move to the chunk which containts element pos
@@ -606,26 +604,26 @@ class ColumnIterator : ChunkingPolicy
606604

607605
mutable unwrap_t<T> const* mCurrent;
608606
int64_t const* mCurrentPos;
607+
uint64_t const* mGlobalOffset;
609608
mutable unwrap_t<T> const* mLast;
610609
arrow::ChunkedArray const* mColumn;
611610
mutable int mFirstIndex;
612611
mutable int mCurrentChunk;
613-
mutable int mOffset;
614612

615613
private:
616614
void checkSkipChunk() const
617615
requires((ChunkingPolicy::chunked == true) && std::same_as<arrow_array_for_t<T>, arrow::ListArray>)
618616
{
619617
auto list = std::static_pointer_cast<arrow::ListArray>(mColumn->chunk(mCurrentChunk));
620-
if (O2_BUILTIN_UNLIKELY(*mCurrentPos - mFirstIndex >= list->length())) {
618+
if (O2_BUILTIN_UNLIKELY(*mCurrentPos + *mGlobalOffset - mFirstIndex >= list->length())) {
621619
nextChunk();
622620
}
623621
}
624622

625623
void checkSkipChunk() const
626624
requires((ChunkingPolicy::chunked == true) && !std::same_as<arrow_array_for_t<T>, arrow::ListArray>)
627625
{
628-
if (O2_BUILTIN_UNLIKELY(((mCurrent + (*mCurrentPos >> SCALE_FACTOR)) >= mLast))) {
626+
if (O2_BUILTIN_UNLIKELY(((mCurrent + ((*mCurrentPos + *mGlobalOffset) >> SCALE_FACTOR)) >= mLast))) {
629627
nextChunk();
630628
}
631629
}
@@ -639,7 +637,6 @@ class ColumnIterator : ChunkingPolicy
639637
requires(std::same_as<arrow_array_for_t<T>, arrow::FixedSizeListArray>)
640638
{
641639
std::shared_ptr<arrow::Array> chunkToUse = mColumn->chunk(mCurrentChunk);
642-
mOffset = chunkToUse->offset();
643640
chunkToUse = std::dynamic_pointer_cast<arrow::FixedSizeListArray>(chunkToUse)->values();
644641
return std::static_pointer_cast<arrow_array_for_t<value_for_t<T>>>(chunkToUse);
645642
}
@@ -648,17 +645,14 @@ class ColumnIterator : ChunkingPolicy
648645
requires(std::same_as<arrow_array_for_t<T>, arrow::ListArray>)
649646
{
650647
std::shared_ptr<arrow::Array> chunkToUse = mColumn->chunk(mCurrentChunk);
651-
mOffset = chunkToUse->offset();
652648
chunkToUse = std::dynamic_pointer_cast<arrow::ListArray>(chunkToUse)->values();
653-
mOffset = chunkToUse->offset();
654649
return std::static_pointer_cast<arrow_array_for_t<value_for_t<T>>>(chunkToUse);
655650
}
656651

657652
auto getCurrentArray() const
658653
requires(!std::same_as<arrow_array_for_t<T>, arrow::FixedSizeListArray> && !std::same_as<arrow_array_for_t<T>, arrow::ListArray>)
659654
{
660655
std::shared_ptr<arrow::Array> chunkToUse = mColumn->chunk(mCurrentChunk);
661-
mOffset = chunkToUse->offset();
662656
return std::static_pointer_cast<arrow_array_for_t<T>>(chunkToUse);
663657
}
664658
};
@@ -1212,7 +1206,7 @@ struct TableIterator : IP, C... {
12121206
{
12131207
using namespace o2::soa;
12141208
auto f = framework::overloaded{
1215-
[this]<soa::is_persistent_column T>(T*) -> void { T::mColumnIterator.mCurrentPos = &this->mRowIndex; },
1209+
[this]<soa::is_persistent_column T>(T*) -> void { T::mColumnIterator.mCurrentPos = &this->mRowIndex; T::mColumnIterator.mGlobalOffset = &this->mOffset; },
12161210
[this]<soa::is_dynamic_column T>(T*) -> void { bindDynamicColumn<T>(typename T::bindings_t{}); },
12171211
[this]<typename T>(T*) -> void {},
12181212
};
@@ -3468,11 +3462,6 @@ struct Concat : Table<o2::aod::Hash<"CONC"_h>, o2::aod::Hash<"CONC/0"_h>, o2::ao
34683462
{
34693463
}
34703464

3471-
Concat(std::vector<std::shared_ptr<arrow::Table>>&& tables)
3472-
: Concat{ArrowHelpers::concatTables(std::move(tables))}
3473-
{
3474-
}
3475-
34763465
Concat(Ts const&... t)
34773466
: Concat{ArrowHelpers::concatTables({t.asArrowTableRef()...})}
34783467
{
@@ -3729,19 +3718,31 @@ class FilteredBase : public T
37293718
}
37303719
}
37313720

3732-
inline void adoptSelection(gandiva::Selection const& selection)
3721+
template <typename S>
3722+
inline void adoptSelection(S)
3723+
{
3724+
}
3725+
3726+
template <typename S>
3727+
requires(std::same_as<std::decay_t<S>, gandiva::Selection>)
3728+
inline void adoptSelection(S selection)
37333729
{
37343730
mSelectedRows = getSpan(selection);
37353731
mCached = false;
37363732
}
37373733

3738-
inline void adoptSelection(SelectionVector&& selection)
3734+
template <typename S>
3735+
requires(std::same_as<std::decay_t<S>, SelectionVector>)
3736+
inline void adoptSelection(S selection)
37393737
{
37403738
mSelectedRowsCache = std::move(selection);
3739+
mSelectedRows = std::span{mSelectedRowsCache};
37413740
mCached = true;
37423741
}
37433742

3744-
inline void adoptSelection(std::span<int64_t const> const& selection)
3743+
template <typename S>
3744+
requires(std::same_as<std::decay_t<S>, std::span<int64_t const>>)
3745+
inline void adoptSelection(S selection)
37453746
{
37463747
mSelectedRows = selection;
37473748
mCached = false;
@@ -3778,7 +3779,7 @@ class Filtered : public FilteredBase<T>
37783779
}
37793780

37803781
Filtered(std::vector<ArrowTableRef>&& tables, is_a_selection auto selection)
3781-
: FilteredBase<T>{std::move(tables), selection} {}
3782+
: FilteredBase<T>{std::move(tables), std::forward<decltype(selection)>(selection)} {}
37823783

37833784
Filtered<T> operator+(is_a_selection auto selection)
37843785
{
@@ -3908,7 +3909,7 @@ class Filtered<Filtered<T>> : public FilteredBase<typename T::table_t>
39083909
}
39093910

39103911
Filtered(std::vector<Filtered<T>>&& tables, is_a_selection auto selection)
3911-
: FilteredBase<typename T::table_t>(std::move(extractTablesFromFiltered(tables)), selection)
3912+
: FilteredBase<typename T::table_t>(std::move(extractTablesFromFiltered(tables)), std::forward<decltype(selection)>(selection))
39123913
{
39133914
for (auto& table : tables) {
39143915
*this *= table;

Framework/Core/src/ASoA.cxx

Lines changed: 11 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -177,42 +177,36 @@ o2::soa::ArrowTableRef ArrowHelpers::concatTables(std::vector<o2::soa::ArrowTabl
177177
return tables.front();
178178
}
179179
std::vector<std::shared_ptr<arrow::ChunkedArray>> columns;
180-
std::vector<std::shared_ptr<arrow::Field>> resultFields = tables[0]->schema()->fields();
180+
std::vector<std::shared_ptr<arrow::Field>> resultFields = tables.front()->schema()->fields();
181181
auto compareFields = [](std::shared_ptr<arrow::Field> const& f1, std::shared_ptr<arrow::Field> const& f2) {
182182
// Let's do this with stable sorting.
183183
return (!f1->Equals(f2)) && (f1->name() < f2->name());
184184
};
185-
std::ranges::for_each(tables.begin() + 1, tables.end(), [&resultFields, &compareFields](auto const& ref) mutable {
186-
std::vector<std::shared_ptr<arrow::Field>> const& fields = ref->fields();
185+
186+
for (auto i = 1; i < tables.size(); ++i) {
187+
auto const& fields = tables[i]->fields();
187188
std::vector<std::shared_ptr<arrow::Field>> intersection;
188189
std::ranges::set_intersection(resultFields, fields, std::back_inserter(intersection), compareFields);
189190
resultFields.swap(intersection);
190-
});
191+
}
191192

192-
std::ranges::transform(resultFields, std::back_inserter(columns), [&tables](auto const& field){
193+
for (auto const& field : resultFields) {
193194
arrow::ArrayVector chunks;
194-
std::ranges::for_each(tables, [&field, &chunks](auto const& table){
195+
for (auto const& table : tables) {
195196
auto ci = table->schema()->GetFieldIndex(field->name());
196197
if (ci == -1) {
197-
throw std::runtime_error("Unable to find field " + field->name());
198+
throw framework::runtime_error_f("Unable to find field {}", field->name().c_str());
198199
}
199200
auto column = table->column(ci);
200201
auto otherChunks = column->chunks();
201202
chunks.insert(chunks.end(), otherChunks.begin(), otherChunks.end());
202-
});
203-
return std::make_shared<arrow::ChunkedArray>(chunks);
204-
});
203+
}
204+
columns.push_back(std::make_shared<arrow::ChunkedArray>(chunks));
205+
}
205206

206207
return {arrow::Table::Make(std::make_shared<arrow::Schema>(resultFields), columns)};
207208
}
208209

209-
o2::soa::ArrowTableRef ArrowHelpers::concatTables(std::vector<std::shared_ptr<arrow::Table>>&& tables)
210-
{
211-
std::vector<ArrowTableRef> refs;
212-
std::ranges::transform(tables, std::back_inserter(refs),[](auto const& table){ return ArrowTableRef{table}; });
213-
return concatTables(std::move(refs));
214-
}
215-
216210
// ASCII-only lowercase. Column labels are plain identifiers, so we deliberately
217211
// avoid the locale-aware std::tolower: it goes through the C locale facet on
218212
// every character and dominated getIndexFromLabel in profiles.

Framework/Core/test/test_ASoA.cxx

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -106,7 +106,9 @@ TEST_CASE("TestTableIteration")
106106

107107
auto i = ColumnIterator<int32_t>(table->column(0).get());
108108
int64_t pos = 0;
109+
uint64_t offset = 0;
109110
i.mCurrentPos = &pos;
111+
i.mGlobalOffset = &offset;
110112
REQUIRE(*i == 0);
111113
pos++;
112114
REQUIRE(*i == 0);
@@ -383,7 +385,7 @@ TEST_CASE("TestConcatTables")
383385
static_assert(std::same_as<NestedJoinTest::columns_t, o2::framework::pack<o2::soa::Index<>, o2::aod::test::Y, o2::aod::test::X, o2::aod::test::Z>>, "Bad nested join");
384386

385387
static_assert(std::same_as<ConcatTest::columns_t, o2::framework::pack<o2::soa::Index<>, o2::aod::test::X>>, "Bad intersection of columns");
386-
ConcatTest tests{tableA, tableB};
388+
ConcatTest tests{{tableA, tableB}};
387389
REQUIRE(16 == tests.size());
388390
for (auto& test : tests) {
389391
REQUIRE(test.index() == test.x());

Framework/Core/test/test_GroupSlicer.cxx

Lines changed: 13 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -195,8 +195,9 @@ TEST_CASE("GroupSlicerSeveralAssociated")
195195
{soa::getLabelFromType<aod::TrksY>(), soa::getMatcherFromTypeForKey<aod::TrksY>(key), key},
196196
{soa::getLabelFromType<aod::TrksZ>(), soa::getMatcherFromTypeForKey<aod::TrksZ>(key), key}});
197197
auto s = slices.updateCacheEntry(0, {trkTableX});
198-
s = slices.updateCacheEntry(1, {trkTableY});
199-
s = slices.updateCacheEntry(2, {trkTableZ});
198+
s &= slices.updateCacheEntry(1, {trkTableY});
199+
s &= slices.updateCacheEntry(2, {trkTableZ});
200+
REQUIRE(s.ok());
200201
o2::framework::GroupSlicer g(e, tt, slices);
201202

202203
auto count = 0;
@@ -631,9 +632,9 @@ TEST_CASE("EmptySliceables")
631632
TEST_CASE("ArrowDirectSlicing")
632633
{
633634
int counts[] = {5, 5, 5, 4, 1};
634-
int offsets[] = {0, 5, 10, 15, 19, 20};
635+
int const offsets[] = {0, 5, 10, 15, 19, 20};
635636
int ids[] = {0, 1, 2, 3, 4};
636-
int sizes[] = {4, 1, 12, 5, 2};
637+
int const sizes[] = {4, 1, 12, 5, 2};
637638

638639
using BigE = soa::Join<aod::Events, aod::EventExtra>;
639640

@@ -683,34 +684,25 @@ TEST_CASE("ArrowDirectSlicing")
683684
REQUIRE(slices_vec[i]->length() == counts[i]);
684685
}
685686

686-
std::vector<arrow::Datum> slices;
687-
std::vector<uint64_t> offsts;
688687
auto bk = Entry(soa::getLabelFromType<aod::Events>(), soa::getMatcherFromTypeForKey<aod::Events>("fID"), "fID");
689688
ArrowTableSlicingCache cache({bk});
690689
auto s = cache.updateCacheEntry(0, {evtTable});
690+
REQUIRE(s.ok());
691691
auto lcache = cache.getCacheFor(bk);
692692
for (auto i = 0u; i < 5; ++i) {
693-
auto [offset, count] = lcache.getSliceFor(i);
694-
auto tbl = b_e.asArrowTableRef().slice({static_cast<uint64_t>(offset), count});
695-
auto ca = tbl->GetColumnByName("fArr");
696-
auto cb = tbl->GetColumnByName("fBoo");
697-
auto cv = tbl->GetColumnByName("fLst");
698-
REQUIRE(ca->length() == counts[i]);
699-
REQUIRE(cb->length() == counts[i]);
700-
REQUIRE(cv->length() == counts[i]);
701-
REQUIRE(ca->Equals(slices_array[i]));
702-
REQUIRE(cb->Equals(slices_bool[i]));
703-
REQUIRE(cv->Equals(slices_vec[i]));
693+
auto [loffset, count] = lcache.getSliceFor(i);
694+
auto tbl = b_e.asArrowTableRef().slice({static_cast<uint64_t>(loffset), count});
695+
REQUIRE(tbl.range.size == counts[i]);
704696
}
705697

706698
int j = 0u;
707699
for (auto i = 0u; i < 5; ++i) {
708-
auto [offset, count] = lcache.getSliceFor(i);
709-
auto tbl = BigE{{b_e.asArrowTableRef().slice({static_cast<uint64_t>(offset), count})}};
700+
auto [loffset, count] = lcache.getSliceFor(i);
701+
auto tbl = BigE{{b_e.asArrowTableRef().slice({static_cast<uint64_t>(loffset), count})}};
710702
REQUIRE(tbl.size() == counts[i]);
711703
for (auto& row : tbl) {
712704
REQUIRE(row.id() == ids[i]);
713-
REQUIRE(row.boo() == (j % 2 == 0));
705+
CHECK(row.boo() == (j % 2 == 0));
714706
auto rid = row.globalIndex();
715707
auto arr = row.arr();
716708
REQUIRE(arr[0] == 0.1f * (float)rid);
@@ -729,7 +721,7 @@ TEST_CASE("ArrowDirectSlicing")
729721

730722
TEST_CASE("TestSlicingException")
731723
{
732-
int offsets[] = {0, 5, 10, 15, 19, 20};
724+
int const offsets[] = {0, 5, 10, 15, 19, 20};
733725
int ids[] = {0, 1, 2, 4, 3};
734726

735727
TableBuilder builderE;

0 commit comments

Comments
 (0)