| 734 | } |
| 735 | |
| 736 | std::vector<int64_t> |
| 737 | HGraph::Add(const DatasetPtr& data) { |
| 738 | std::vector<int64_t> failed_ids; |
| 739 | auto base_dim = data->GetDim(); |
| 740 | if (data_type_ != DataTypes::DATA_TYPE_SPARSE) { |
| 741 | CHECK_ARGUMENT(base_dim == dim_, |
| 742 | fmt::format("base.dim({}) must be equal to index.dim({})", base_dim, dim_)); |
| 743 | } |
| 744 | CHECK_ARGUMENT(get_data(data) != nullptr, "base.float_vector is nullptr"); |
| 745 | |
| 746 | { |
| 747 | std::scoped_lock lock(this->add_mutex_); |
| 748 | if (this->total_count_ == 0) { |
| 749 | this->Train(data); |
| 750 | } |
| 751 | } |
| 752 | |
| 753 | auto add_func = [&](const void* data, |
| 754 | int level, |
| 755 | InnerIdType inner_id, |
| 756 | const char* extra_info, |
| 757 | const AttributeSet* attrs) -> void { |
| 758 | if (this->extra_infos_ != nullptr) { |
| 759 | this->extra_infos_->InsertExtraInfo(extra_info, inner_id); |
| 760 | } |
| 761 | if (attrs != nullptr and this->use_attribute_filter_) { |
| 762 | this->attr_filter_index_->Insert(*attrs, inner_id); |
| 763 | } |
| 764 | this->add_one_point(data, level, inner_id); |
| 765 | }; |
| 766 | |
| 767 | std::vector<std::future<void>> futures; |
| 768 | auto total = data->GetNumElements(); |
| 769 | const auto* labels = data->GetIds(); |
| 770 | const auto* extra_infos = data->GetExtraInfos(); |
| 771 | const auto* attr_sets = data->GetAttributeSets(); |
| 772 | Vector<std::pair<InnerIdType, LabelType>> inner_ids(allocator_); |
| 773 | for (int64_t j = 0; j < total; ++j) { |
| 774 | InnerIdType inner_id; |
| 775 | |
| 776 | // try recover tombstone |
| 777 | if (this->data_type_ != DataTypes::DATA_TYPE_SPARSE) { |
| 778 | auto one_base = get_single_dataset(data, j); |
| 779 | bool is_process_finished = try_recover_tombstone(one_base, failed_ids); |
| 780 | if (is_process_finished) { |
| 781 | continue; |
| 782 | } |
| 783 | } |
| 784 | |
| 785 | { |
| 786 | std::scoped_lock lock(this->add_mutex_); |
| 787 | inner_id = this->get_unique_inner_ids(1).at(0); |
| 788 | this->resize(total_count_.load() + 1); |
| 789 | ++total_count_; |
| 790 | } |
| 791 | |
| 792 | { |
| 793 | std::scoped_lock label_lock(this->label_lookup_mutex_); |
no test coverage detected