| 68 | feat_desc_(feat_desc) {} |
| 69 | |
| 70 | Status Init(const Tensor& default_tensor, int64 default_value_dim) { |
| 71 | if (storage_ == nullptr) { |
| 72 | return errors::InvalidArgument( |
| 73 | "Invalid ht_type to construct EmbeddingVar"); |
| 74 | } |
| 75 | |
| 76 | storage_type_ = storage_->GetStorageType(); |
| 77 | filter_ = FilterFactory::CreateFilter<K, V, EmbeddingVar<K, V>>( |
| 78 | emb_config_, this, storage_, feat_desc_); |
| 79 | emb_config_.default_value_dim = default_value_dim; |
| 80 | value_len_ = |
| 81 | default_tensor.NumElements() / emb_config_.default_value_dim; |
| 82 | |
| 83 | if (storage_->IsUseHbm()) { |
| 84 | #if GOOGLE_CUDA |
| 85 | default_value_ = TypedAllocator::Allocate<V>(alloc_, |
| 86 | default_tensor.NumElements(), AllocationAttributes()); |
| 87 | auto default_tensor_flat = default_tensor.flat<V>(); |
| 88 | dev_addr_buffer_ = nullptr; |
| 89 | dev_addr_buffer_size_ = 0; |
| 90 | cudaMemcpy(default_value_, &default_tensor_flat(0), |
| 91 | default_tensor.TotalBytes(), cudaMemcpyDeviceToDevice); |
| 92 | #endif // GOOGLE_CUDA |
| 93 | } else if (storage_->IsSingleHbm()) { |
| 94 | #if GOOGLE_CUDA |
| 95 | storage_->SetValueLen(value_len_); |
| 96 | default_value_ = TypedAllocator::Allocate<V>( |
| 97 | alloc_, default_tensor.NumElements(), AllocationAttributes()); |
| 98 | auto default_tensor_flat = default_tensor.flat<V>(); |
| 99 | cudaMemcpy(default_value_, &default_tensor_flat(0), |
| 100 | default_tensor.TotalBytes(), cudaMemcpyDeviceToDevice); |
| 101 | #endif // GOOGLE_CUDA |
| 102 | } else { |
| 103 | alloc_ = ev_allocator(); |
| 104 | default_value_ = TypedAllocator::Allocate<V>(default_value_alloc_, |
| 105 | default_tensor.NumElements(), AllocationAttributes()); |
| 106 | |
| 107 | auto default_tensor_flat = default_tensor.flat<V>(); |
| 108 | memcpy(default_value_, &default_tensor_flat(0), |
| 109 | default_tensor.TotalBytes()); |
| 110 | |
| 111 | default_value_no_permission_ = TypedAllocator::Allocate<V>( |
| 112 | default_value_alloc_, value_len_, AllocationAttributes()); |
| 113 | for (int i = 0; i < value_len_; ++i) { |
| 114 | default_value_no_permission_[i] = static_cast<V>( |
| 115 | emb_config_.default_value_no_permission); |
| 116 | } |
| 117 | } |
| 118 | bool is_all_slots_initialized = |
| 119 | feat_desc_->InitSlotInfo( |
| 120 | emb_config_.emb_index, value_len_, |
| 121 | std::pair<V*, int64>( |
| 122 | default_value_, emb_config_.default_value_dim)); |
| 123 | if (is_all_slots_initialized) { |
| 124 | storage_->Init(); |
| 125 | } |
| 126 | |
| 127 | return Status::OK(); |
nothing calls this directly
no test coverage detected