MCPcopy Create free account
hub / github.com/apache/singa / Data

Class Data

examples/rnn/char_rnn.py:92–121  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

90
91
92class Data(object):
93
94 def __init__(self, fpath, batch_size=32, seq_length=100, train_ratio=0.8):
95 '''Data object for loading a plain text file.
96
97 Args:
98 fpath, path to the text file.
99 train_ratio, split the text file into train and test sets, where
100 train_ratio of the characters are in the train set.
101 '''
102 self.raw_data = open(fpath, 'r',
103 encoding='iso-8859-1').read() # read text file
104 chars = list(set(self.raw_data))
105 self.vocab_size = len(chars)
106 self.char_to_idx = {ch: i for i, ch in enumerate(chars)}
107 self.idx_to_char = {i: ch for i, ch in enumerate(chars)}
108 data = [self.char_to_idx[c] for c in self.raw_data]
109 # seq_length + 1 for the data + label
110 nsamples = len(data) // (1 + seq_length)
111 data = data[0:nsamples * (1 + seq_length)]
112 data = np.asarray(data, dtype=np.int32)
113 data = np.reshape(data, (-1, seq_length + 1))
114 # shuffle all sequences
115 np.random.shuffle(data)
116 self.train_dat = data[0:int(data.shape[0] * train_ratio)]
117 self.num_train_batch = self.train_dat.shape[0] // batch_size
118 self.val_dat = data[self.train_dat.shape[0]:]
119 self.num_test_batch = self.val_dat.shape[0] // batch_size
120 print('train dat', self.train_dat.shape)
121 print('val dat', self.val_dat.shape)
122
123
124def numpy2tensors(npx, npy, dev, inputs=None, labels=None):

Callers 1

char_rnn.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected