MCPcopy Create free account
hub / github.com/bigdata-ustc/Agent4Edu / ValTestDataLoader

Class ValTestDataLoader

Code/tools/irt/original/data_loader.py:51–94  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

49
50
51class 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

Callers 1

validateFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected