| 49 | |
| 50 | |
| 51 | class ValTestDataLoader(object): |
| 52 | def __init__(self, d_type='validation'): |
| 53 | self.ptr = 0 |
| 54 | self.data = [] |
| 55 | self.d_type = d_type |
| 56 | |
| 57 | if d_type == 'validation': |
| 58 | data_file = 'data/val_set.json' |
| 59 | else: |
| 60 | data_file = 'data/test_set.json' |
| 61 | config_file = 'config.txt' |
| 62 | with open(data_file, encoding='utf8') as i_f: |
| 63 | self.data = json.load(i_f) |
| 64 | with open(config_file) as i_f: |
| 65 | i_f.readline() |
| 66 | _, _, knowledge_n = i_f.readline().split(',') |
| 67 | self.knowledge_dim = int(knowledge_n) |
| 68 | |
| 69 | def next_batch(self): |
| 70 | if self.is_end(): |
| 71 | return None, None, None, None |
| 72 | logs = self.data[self.ptr]['logs'] |
| 73 | user_id = self.data[self.ptr]['user_id'] |
| 74 | input_stu_ids, input_exer_ids, input_knowledge_embs, ys = [], [], [], [] |
| 75 | for log in logs: |
| 76 | input_stu_ids.append(user_id) |
| 77 | input_exer_ids.append(log['exer_id']) |
| 78 | knowledge_emb = [0.] * self.knowledge_dim |
| 79 | for knowledge_code in [log['knowledge_code']]: |
| 80 | knowledge_emb[knowledge_code] = 1.0 |
| 81 | input_knowledge_embs.append(knowledge_emb) |
| 82 | y = log['score'] |
| 83 | ys.append(y) |
| 84 | self.ptr += 1 |
| 85 | return torch.LongTensor(input_stu_ids), torch.LongTensor(input_exer_ids), torch.Tensor(input_knowledge_embs), torch.LongTensor(ys) |
| 86 | |
| 87 | def is_end(self): |
| 88 | if self.ptr >= len(self.data): |
| 89 | return True |
| 90 | else: |
| 91 | return False |
| 92 | |
| 93 | def reset(self): |
| 94 | self.ptr = 0 |