tprism.embedding_generator

This module contains EmbeddingGenerators, which assign a tensor to an atom.

  1"""
  2This module contains EmbeddingGenerators,
  3which assign a tensor to an atom.
  4"""
  5import torch
  6import json
  7import logging
  8import numpy as np
  9
 10import os
 11import h5py
 12
 13from tprism.util import TensorInfoMapper, debug_logger
 14from tprism.placeholder import PlaceholderData
 15from numpy import ndarray
 16from torch import Tensor
 17from typing import Any, Dict, Optional, Sequence, Tuple, TypedDict
 18
 19logger = logging.getLogger(__name__)
 20embedding_logger = debug_logger("embedding")
 21feed_logger = debug_logger("feed")
 22
 23EmbeddingData=Dict[str,ndarray]
 24
 25
 26class CycleEmbeddingEntry(TypedDict):
 27    tensor: PlaceholderData
 28    data: Tensor
 29    id: int
 30
 31def load_embedding_data(filename: str, key: str) -> EmbeddingData:
 32    """Load input data supporting .h5/.json format
 33
 34    Args:
 35        data_filename_list: list of input file names
 36    Returns:
 37        merged input data
 38
 39    """
 40    _, ext = os.path.splitext(filename)
 41    dataset:EmbeddingData = {}
 42    if ext == ".h5":
 43        logger.info("[LOAD] %s", filename)
 44        dataset = load_embedding_h5(filename, key)
 45    elif ext == ".json":
 46        logger.info("[LOAD] %s", filename)
 47        dataset = load_embedding_npy(filename, key)
 48    else:
 49        logger.error("unknown embedding format: %s", filename)
 50    return dataset
 51
 52def load_embedding_h5(filename, key)-> EmbeddingData:
 53    infh:Any = h5py.File(filename, "r")
 54    dataset:EmbeddingData = {}
 55    if key in infh:
 56        for vocab_name in infh[key]:
 57            rs = infh[key][vocab_name][()]
 58            dataset[vocab_name] = rs
 59            embedding_logger.debug("[LOAD DatasetEmbedding] %s", vocab_name)
 60    infh.close()
 61    return dataset
 62
 63
 64def load_embedding_npy(filename: str, key: str) -> EmbeddingData:
 65    fp = open(filename, "r")
 66    obj = json.load(fp)
 67    dataset:EmbeddingData = {}
 68    if key in obj["group"]:
 69        tensor=np.load(obj["filename"])
 70        name=obj["name"]
 71        dataset[name]=tensor
 72        embedding_logger.debug("[LOAD DatasetEmbedding] %s", name)
 73    return dataset
 74
 75
 76class BaseEmbeddingGenerator:
 77    def __init__(self):
 78        pass
 79
 80    def is_embedding(self, vocab_name: str) -> bool:
 81        return False
 82
 83    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
 84        return ()
 85
 86    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id = None) -> Optional[PlaceholderData]:
 87        return None
 88
 89    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
 90        return torch.tensor(0.0)
 91    
 92    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
 93        return torch.tensor(0.0)
 94    
 95    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
 96        return feed_dict
 97
 98
 99class CycleEmbeddingGenerator(BaseEmbeddingGenerator):
100    def __init__(self):
101        super().__init__()
102        self.embedding: Dict[str, CycleEmbeddingEntry] = {}
103        self.tensor_shape: TensorInfoMapper = TensorInfoMapper()
104
105    def load(self, tensor_shape: TensorInfoMapper) -> None:
106        self.tensor_shape = tensor_shape
107
108    def get_embedding(
109        self,
110        vocab_name: str,
111        shape: Optional[Tuple[int, ...]] = None,
112        node_id: Optional[int] = None,
113    ) -> Optional[PlaceholderData]:
114        return None
115
116    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
117        ph_name: str = name + "_cyc"
118        shape_tuple: Tuple[int, ...] = tuple(shape)
119        if ph_name in self.embedding:
120            embedding_logger.debug("[GET cycle]> %s : %s", ph_name, self.embedding[ph_name]["tensor"])
121            return torch.tensor(self.embedding[ph_name]["data"])
122        else:
123            embedding_logger.debug("[CREATE cycle]> %s : %s", ph_name, shape_tuple)
124            self.embedding[ph_name] = {
125                "tensor": PlaceholderData(
126                    name=ph_name, shape=shape_tuple, dtype=torch.float32
127                ),
128                "data": torch.tensor(np.zeros(shape=shape_tuple, dtype=np.float32)),
129                "id": node_id,
130            }
131            return torch.tensor(self.embedding[ph_name]["data"])
132
133    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]: ## idx is not used
134        for ph_name, data in self.embedding.items():
135            batch_data = data["data"]
136            ph_var = data["tensor"]
137            feed_logger.debug("[cycle feed] node_id: %s => %s", data["id"], ph_name)
138            feed_dict[ph_var] = torch.Tensor(batch_data)
139        return feed_dict
140
141    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
142        total_loss: Tensor = torch.tensor(0.0)
143        for ph_name, data in self.embedding.items():
144            node_id: int = data["id"]
145            embedding_logger.debug("[cycle update] node_id: %s => %s", node_id, ph_name)
146            ##
147            o = out_inside[node_id]
148            loss = self.embedding[ph_name]["data"] - o
149            total_loss += (loss ** 2).sum()
150            ##
151            self.embedding[ph_name]["data"] = o
152            # a=0.5
153            # self.embedding[ph_name]["data"]=(1.0-a)*self.embedding[ph_name]["data"]+a*out_inside[node_id]
154        return total_loss
155
156
157# embedding data from data
158class EmbeddingGenerator(BaseEmbeddingGenerator):
159    """ Generating embedding data from the given tensor atom
160    
161        1. If a given tensor atom is not shared by this generator, do nothing
162        2. A placeholder name is computed by the given tensor atom
163        3. If the placeholder already exists, return this placeholder.
164        4. Create a placeholder and return it.
165
166    Attributes:
167        dataset (Dict[str,tensor]) : loaded dataset which is defined by dictionary from a tensor atom name to an assigned tensor.
168        created_ph_var (Dict[str,PlaceholderData]) : assign
169    """
170    def __init__(self, const_flag:bool=False) -> None:
171        super().__init__()
172        self.dataset: EmbeddingData = {}
173        self.created_ph_var: Dict[str, PlaceholderData] = {}
174        self.const_flag=const_flag
175
176    def load(self, filename: str, key: str="train") -> None:
177        self.dataset=load_embedding_data(filename, key)
178
179    def is_embedding(self, vocab_name: str) -> bool:
180        return vocab_name in self.dataset
181
182    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
183        return self.dataset[vocab_name].shape
184
185    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] =None, node_id = None) -> Optional[PlaceholderData]:
186        if not self.is_embedding(vocab_name):
187            embedding_logger.debug("[SKIP]> %s", vocab_name)
188            return None
189        ph_name = vocab_name + "_ph"
190        if ph_name in self.created_ph_var:
191            embedding_logger.debug("[GET]> %s : %s", ph_name, self.created_ph_var[ph_name])
192            return self.created_ph_var[ph_name]
193        else:
194            if shape is None:
195                shape = self.dataset[vocab_name].shape
196            if self.const_flag:
197                self.created_ph_var[ph_name] = PlaceholderData(
198                    name=ph_name, shape=shape, dtype=torch.float32, ref=vocab_name
199                )
200                embedding_logger.debug("[CREATE const]> %s : %s", ph_name, shape)
201            else:
202                self.created_ph_var[ph_name] = PlaceholderData(
203                    name=ph_name, shape=shape, dtype=torch.float32
204                )
205                embedding_logger.debug("[CREATE]> %s : %s ref: %s", ph_name, shape, vocab_name)
206            return self.created_ph_var[ph_name]
207        return None
208    
209    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
210        for vocab_name, data in self.dataset.items():
211            ph_name = vocab_name + "_ph"
212            if idx is None or self.const_flag:
213                batch_data = data
214            else:
215                batch_data = data[idx]
216            if ph_name in self.created_ph_var:
217                ph_var = self.created_ph_var[ph_name]
218                feed_dict[ph_var] = torch.Tensor(batch_data)
219            feed_logger.debug("[feed] %s => %s", vocab_name, ph_name)
220        return feed_dict
logger = <Logger tprism.embedding_generator (INFO)>
embedding_logger = <Logger tprism.debug.embedding (INFO)>
feed_logger = <Logger tprism.debug.feed (INFO)>
EmbeddingData = typing.Dict[str, numpy.ndarray]
class CycleEmbeddingEntry(typing.TypedDict):
27class CycleEmbeddingEntry(TypedDict):
28    tensor: PlaceholderData
29    data: Tensor
30    id: int
data: torch.Tensor
id: int
def load_embedding_data(filename: str, key: str) -> Dict[str, numpy.ndarray]:
32def load_embedding_data(filename: str, key: str) -> EmbeddingData:
33    """Load input data supporting .h5/.json format
34
35    Args:
36        data_filename_list: list of input file names
37    Returns:
38        merged input data
39
40    """
41    _, ext = os.path.splitext(filename)
42    dataset:EmbeddingData = {}
43    if ext == ".h5":
44        logger.info("[LOAD] %s", filename)
45        dataset = load_embedding_h5(filename, key)
46    elif ext == ".json":
47        logger.info("[LOAD] %s", filename)
48        dataset = load_embedding_npy(filename, key)
49    else:
50        logger.error("unknown embedding format: %s", filename)
51    return dataset

Load input data supporting .h5/.json format

Arguments:
  • data_filename_list: list of input file names
Returns:

merged input data

def load_embedding_h5(filename, key) -> Dict[str, numpy.ndarray]:
53def load_embedding_h5(filename, key)-> EmbeddingData:
54    infh:Any = h5py.File(filename, "r")
55    dataset:EmbeddingData = {}
56    if key in infh:
57        for vocab_name in infh[key]:
58            rs = infh[key][vocab_name][()]
59            dataset[vocab_name] = rs
60            embedding_logger.debug("[LOAD DatasetEmbedding] %s", vocab_name)
61    infh.close()
62    return dataset
def load_embedding_npy(filename: str, key: str) -> Dict[str, numpy.ndarray]:
65def load_embedding_npy(filename: str, key: str) -> EmbeddingData:
66    fp = open(filename, "r")
67    obj = json.load(fp)
68    dataset:EmbeddingData = {}
69    if key in obj["group"]:
70        tensor=np.load(obj["filename"])
71        name=obj["name"]
72        dataset[name]=tensor
73        embedding_logger.debug("[LOAD DatasetEmbedding] %s", name)
74    return dataset
class BaseEmbeddingGenerator:
77class BaseEmbeddingGenerator:
78    def __init__(self):
79        pass
80
81    def is_embedding(self, vocab_name: str) -> bool:
82        return False
83
84    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
85        return ()
86
87    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id = None) -> Optional[PlaceholderData]:
88        return None
89
90    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
91        return torch.tensor(0.0)
92    
93    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
94        return torch.tensor(0.0)
95    
96    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
97        return feed_dict
def is_embedding(self, vocab_name: str) -> bool:
81    def is_embedding(self, vocab_name: str) -> bool:
82        return False
def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
84    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
85        return ()
def get_embedding( self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id=None) -> Optional[tprism.placeholder.PlaceholderData]:
87    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id = None) -> Optional[PlaceholderData]:
88        return None
def update(self, out_inside: Sequence[torch.Tensor]) -> torch.Tensor:
90    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
91        return torch.tensor(0.0)
def forward(self, name: str, shape: Sequence[int], node_id: int) -> torch.Tensor:
93    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
94        return torch.tensor(0.0)
def build_feed( self, feed_dict: Dict[tprism.placeholder.PlaceholderData, torch.Tensor], idx: Optional[numpy.ndarray] = None) -> Dict[tprism.placeholder.PlaceholderData, torch.Tensor]:
96    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
97        return feed_dict
class CycleEmbeddingGenerator(BaseEmbeddingGenerator):
100class CycleEmbeddingGenerator(BaseEmbeddingGenerator):
101    def __init__(self):
102        super().__init__()
103        self.embedding: Dict[str, CycleEmbeddingEntry] = {}
104        self.tensor_shape: TensorInfoMapper = TensorInfoMapper()
105
106    def load(self, tensor_shape: TensorInfoMapper) -> None:
107        self.tensor_shape = tensor_shape
108
109    def get_embedding(
110        self,
111        vocab_name: str,
112        shape: Optional[Tuple[int, ...]] = None,
113        node_id: Optional[int] = None,
114    ) -> Optional[PlaceholderData]:
115        return None
116
117    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
118        ph_name: str = name + "_cyc"
119        shape_tuple: Tuple[int, ...] = tuple(shape)
120        if ph_name in self.embedding:
121            embedding_logger.debug("[GET cycle]> %s : %s", ph_name, self.embedding[ph_name]["tensor"])
122            return torch.tensor(self.embedding[ph_name]["data"])
123        else:
124            embedding_logger.debug("[CREATE cycle]> %s : %s", ph_name, shape_tuple)
125            self.embedding[ph_name] = {
126                "tensor": PlaceholderData(
127                    name=ph_name, shape=shape_tuple, dtype=torch.float32
128                ),
129                "data": torch.tensor(np.zeros(shape=shape_tuple, dtype=np.float32)),
130                "id": node_id,
131            }
132            return torch.tensor(self.embedding[ph_name]["data"])
133
134    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]: ## idx is not used
135        for ph_name, data in self.embedding.items():
136            batch_data = data["data"]
137            ph_var = data["tensor"]
138            feed_logger.debug("[cycle feed] node_id: %s => %s", data["id"], ph_name)
139            feed_dict[ph_var] = torch.Tensor(batch_data)
140        return feed_dict
141
142    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
143        total_loss: Tensor = torch.tensor(0.0)
144        for ph_name, data in self.embedding.items():
145            node_id: int = data["id"]
146            embedding_logger.debug("[cycle update] node_id: %s => %s", node_id, ph_name)
147            ##
148            o = out_inside[node_id]
149            loss = self.embedding[ph_name]["data"] - o
150            total_loss += (loss ** 2).sum()
151            ##
152            self.embedding[ph_name]["data"] = o
153            # a=0.5
154            # self.embedding[ph_name]["data"]=(1.0-a)*self.embedding[ph_name]["data"]+a*out_inside[node_id]
155        return total_loss
embedding: Dict[str, CycleEmbeddingEntry]
def load(self, tensor_shape: tprism.util.TensorInfoMapper) -> None:
106    def load(self, tensor_shape: TensorInfoMapper) -> None:
107        self.tensor_shape = tensor_shape
def get_embedding( self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id: Optional[int] = None) -> Optional[tprism.placeholder.PlaceholderData]:
109    def get_embedding(
110        self,
111        vocab_name: str,
112        shape: Optional[Tuple[int, ...]] = None,
113        node_id: Optional[int] = None,
114    ) -> Optional[PlaceholderData]:
115        return None
def forward(self, name: str, shape: Sequence[int], node_id: int) -> torch.Tensor:
117    def forward(self, name: str, shape: Sequence[int], node_id: int) -> Tensor:
118        ph_name: str = name + "_cyc"
119        shape_tuple: Tuple[int, ...] = tuple(shape)
120        if ph_name in self.embedding:
121            embedding_logger.debug("[GET cycle]> %s : %s", ph_name, self.embedding[ph_name]["tensor"])
122            return torch.tensor(self.embedding[ph_name]["data"])
123        else:
124            embedding_logger.debug("[CREATE cycle]> %s : %s", ph_name, shape_tuple)
125            self.embedding[ph_name] = {
126                "tensor": PlaceholderData(
127                    name=ph_name, shape=shape_tuple, dtype=torch.float32
128                ),
129                "data": torch.tensor(np.zeros(shape=shape_tuple, dtype=np.float32)),
130                "id": node_id,
131            }
132            return torch.tensor(self.embedding[ph_name]["data"])
def build_feed( self, feed_dict: Dict[tprism.placeholder.PlaceholderData, torch.Tensor], idx: Optional[numpy.ndarray] = None) -> Dict[tprism.placeholder.PlaceholderData, torch.Tensor]:
134    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]: ## idx is not used
135        for ph_name, data in self.embedding.items():
136            batch_data = data["data"]
137            ph_var = data["tensor"]
138            feed_logger.debug("[cycle feed] node_id: %s => %s", data["id"], ph_name)
139            feed_dict[ph_var] = torch.Tensor(batch_data)
140        return feed_dict
def update(self, out_inside: Sequence[torch.Tensor]) -> torch.Tensor:
142    def update(self, out_inside: Sequence[Tensor]) -> Tensor:
143        total_loss: Tensor = torch.tensor(0.0)
144        for ph_name, data in self.embedding.items():
145            node_id: int = data["id"]
146            embedding_logger.debug("[cycle update] node_id: %s => %s", node_id, ph_name)
147            ##
148            o = out_inside[node_id]
149            loss = self.embedding[ph_name]["data"] - o
150            total_loss += (loss ** 2).sum()
151            ##
152            self.embedding[ph_name]["data"] = o
153            # a=0.5
154            # self.embedding[ph_name]["data"]=(1.0-a)*self.embedding[ph_name]["data"]+a*out_inside[node_id]
155        return total_loss
class EmbeddingGenerator(BaseEmbeddingGenerator):
159class EmbeddingGenerator(BaseEmbeddingGenerator):
160    """ Generating embedding data from the given tensor atom
161    
162        1. If a given tensor atom is not shared by this generator, do nothing
163        2. A placeholder name is computed by the given tensor atom
164        3. If the placeholder already exists, return this placeholder.
165        4. Create a placeholder and return it.
166
167    Attributes:
168        dataset (Dict[str,tensor]) : loaded dataset which is defined by dictionary from a tensor atom name to an assigned tensor.
169        created_ph_var (Dict[str,PlaceholderData]) : assign
170    """
171    def __init__(self, const_flag:bool=False) -> None:
172        super().__init__()
173        self.dataset: EmbeddingData = {}
174        self.created_ph_var: Dict[str, PlaceholderData] = {}
175        self.const_flag=const_flag
176
177    def load(self, filename: str, key: str="train") -> None:
178        self.dataset=load_embedding_data(filename, key)
179
180    def is_embedding(self, vocab_name: str) -> bool:
181        return vocab_name in self.dataset
182
183    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
184        return self.dataset[vocab_name].shape
185
186    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] =None, node_id = None) -> Optional[PlaceholderData]:
187        if not self.is_embedding(vocab_name):
188            embedding_logger.debug("[SKIP]> %s", vocab_name)
189            return None
190        ph_name = vocab_name + "_ph"
191        if ph_name in self.created_ph_var:
192            embedding_logger.debug("[GET]> %s : %s", ph_name, self.created_ph_var[ph_name])
193            return self.created_ph_var[ph_name]
194        else:
195            if shape is None:
196                shape = self.dataset[vocab_name].shape
197            if self.const_flag:
198                self.created_ph_var[ph_name] = PlaceholderData(
199                    name=ph_name, shape=shape, dtype=torch.float32, ref=vocab_name
200                )
201                embedding_logger.debug("[CREATE const]> %s : %s", ph_name, shape)
202            else:
203                self.created_ph_var[ph_name] = PlaceholderData(
204                    name=ph_name, shape=shape, dtype=torch.float32
205                )
206                embedding_logger.debug("[CREATE]> %s : %s ref: %s", ph_name, shape, vocab_name)
207            return self.created_ph_var[ph_name]
208        return None
209    
210    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
211        for vocab_name, data in self.dataset.items():
212            ph_name = vocab_name + "_ph"
213            if idx is None or self.const_flag:
214                batch_data = data
215            else:
216                batch_data = data[idx]
217            if ph_name in self.created_ph_var:
218                ph_var = self.created_ph_var[ph_name]
219                feed_dict[ph_var] = torch.Tensor(batch_data)
220            feed_logger.debug("[feed] %s => %s", vocab_name, ph_name)
221        return feed_dict

Generating embedding data from the given tensor atom

1. If a given tensor atom is not shared by this generator, do nothing
2. A placeholder name is computed by the given tensor atom
3. If the placeholder already exists, return this placeholder.
4. Create a placeholder and return it.
Attributes:
  • dataset (Dict[str,tensor]) : loaded dataset which is defined by dictionary from a tensor atom name to an assigned tensor.
  • created_ph_var (Dict[str,PlaceholderData]) : assign
EmbeddingGenerator(const_flag: bool = False)
171    def __init__(self, const_flag:bool=False) -> None:
172        super().__init__()
173        self.dataset: EmbeddingData = {}
174        self.created_ph_var: Dict[str, PlaceholderData] = {}
175        self.const_flag=const_flag
dataset: Dict[str, numpy.ndarray]
created_ph_var: Dict[str, tprism.placeholder.PlaceholderData]
const_flag
def load(self, filename: str, key: str = 'train') -> None:
177    def load(self, filename: str, key: str="train") -> None:
178        self.dataset=load_embedding_data(filename, key)
def is_embedding(self, vocab_name: str) -> bool:
180    def is_embedding(self, vocab_name: str) -> bool:
181        return vocab_name in self.dataset
def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
183    def get_shape(self, vocab_name: str) -> Tuple[int, ...]:
184        return self.dataset[vocab_name].shape
def get_embedding( self, vocab_name: str, shape: Optional[Tuple[int, ...]] = None, node_id=None) -> Optional[tprism.placeholder.PlaceholderData]:
186    def get_embedding(self, vocab_name: str, shape: Optional[Tuple[int, ...]] =None, node_id = None) -> Optional[PlaceholderData]:
187        if not self.is_embedding(vocab_name):
188            embedding_logger.debug("[SKIP]> %s", vocab_name)
189            return None
190        ph_name = vocab_name + "_ph"
191        if ph_name in self.created_ph_var:
192            embedding_logger.debug("[GET]> %s : %s", ph_name, self.created_ph_var[ph_name])
193            return self.created_ph_var[ph_name]
194        else:
195            if shape is None:
196                shape = self.dataset[vocab_name].shape
197            if self.const_flag:
198                self.created_ph_var[ph_name] = PlaceholderData(
199                    name=ph_name, shape=shape, dtype=torch.float32, ref=vocab_name
200                )
201                embedding_logger.debug("[CREATE const]> %s : %s", ph_name, shape)
202            else:
203                self.created_ph_var[ph_name] = PlaceholderData(
204                    name=ph_name, shape=shape, dtype=torch.float32
205                )
206                embedding_logger.debug("[CREATE]> %s : %s ref: %s", ph_name, shape, vocab_name)
207            return self.created_ph_var[ph_name]
208        return None
def build_feed( self, feed_dict: Dict[tprism.placeholder.PlaceholderData, torch.Tensor], idx: Optional[numpy.ndarray] = None) -> Dict[tprism.placeholder.PlaceholderData, torch.Tensor]:
210    def build_feed(self, feed_dict: Dict[PlaceholderData, Tensor], idx: Optional[ndarray]=None) -> Dict[PlaceholderData, Tensor]:
211        for vocab_name, data in self.dataset.items():
212            ph_name = vocab_name + "_ph"
213            if idx is None or self.const_flag:
214                batch_data = data
215            else:
216                batch_data = data[idx]
217            if ph_name in self.created_ph_var:
218                ph_var = self.created_ph_var[ph_name]
219                feed_dict[ph_var] = torch.Tensor(batch_data)
220            feed_logger.debug("[feed] %s => %s", vocab_name, ph_name)
221        return feed_dict