MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Init

Method Init

tensorflow/core/framework/embedding/embedding_var.h:70–128  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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();

Callers

nothing calls this directly

Calls 10

InvalidArgumentFunction · 0.85
ev_allocatorFunction · 0.85
GetStorageTypeMethod · 0.80
NumElementsMethod · 0.45
IsUseHbmMethod · 0.45
TotalBytesMethod · 0.45
IsSingleHbmMethod · 0.45
SetValueLenMethod · 0.45
InitSlotInfoMethod · 0.45

Tested by

no test coverage detected