| 54 | |
| 55 | |
| 56 | class PadDataset(Dataset): |
| 57 | def __init__(self, dataset, seq_len, eod_id): |
| 58 | self.dataset = dataset |
| 59 | self.seq_len = seq_len + 1 |
| 60 | self.eod_id = eod_id |
| 61 | |
| 62 | def __len__(self): |
| 63 | return len(self.dataset) |
| 64 | |
| 65 | def __getitem__(self, idx): |
| 66 | item = self.dataset[idx][0] |
| 67 | return (item[:self.seq_len],) if self.seq_len <= len(item) else ( |
| 68 | np.concatenate((item, np.ones(self.seq_len - len(item)) * self.eod_id), axis=0),) |
| 69 | # return (np.pad(item, (0, 1), constant_values=self.eod_id),) |
| 70 | |
| 71 | |
| 72 | class BinaryDataset(Dataset): |
no outgoing calls
no test coverage detected