1use crate::{ModelConfig, ModelStats, Triple};
4use anyhow::Result;
5use chrono::{DateTime, Utc};
6#[allow(unused_imports)]
7use scirs2_core::random::{Random, RngExt};
8use std::collections::{HashMap, HashSet};
9use uuid::Uuid;
10
11#[derive(Debug, Clone)]
13pub struct BaseModel {
14 pub config: ModelConfig,
16 pub model_id: Uuid,
18 pub entity_to_id: HashMap<String, usize>,
20 pub id_to_entity: HashMap<usize, String>,
22 pub relation_to_id: HashMap<String, usize>,
24 pub id_to_relation: HashMap<usize, String>,
26 pub triples: Vec<(usize, usize, usize)>,
28 pub positive_triples: HashSet<(usize, usize, usize)>,
30 pub is_trained: bool,
32 pub creation_time: DateTime<Utc>,
34 pub last_training_time: Option<DateTime<Utc>>,
36}
37
38impl BaseModel {
39 pub fn new(config: ModelConfig) -> Self {
41 Self {
42 model_id: Uuid::new_v4(),
43 config,
44 entity_to_id: HashMap::new(),
45 id_to_entity: HashMap::new(),
46 relation_to_id: HashMap::new(),
47 id_to_relation: HashMap::new(),
48 triples: Vec::new(),
49 positive_triples: HashSet::new(),
50 is_trained: false,
51 creation_time: Utc::now(),
52 last_training_time: None,
53 }
54 }
55
56 pub fn add_triple(&mut self, triple: Triple) -> Result<()> {
58 let subject_str = triple.subject.to_string();
59 let predicate_str = triple.predicate.to_string();
60 let object_str = triple.object.to_string();
61
62 let subject_id = self.get_or_create_entity_id(subject_str);
64 let object_id = self.get_or_create_entity_id(object_str);
65
66 let predicate_id = self.get_or_create_relation_id(predicate_str);
68
69 let triple_ids = (subject_id, predicate_id, object_id);
71 if !self.positive_triples.contains(&triple_ids) {
72 self.triples.push(triple_ids);
73 self.positive_triples.insert(triple_ids);
74 }
75
76 Ok(())
77 }
78
79 fn get_or_create_entity_id(&mut self, entity: String) -> usize {
81 if let Some(&id) = self.entity_to_id.get(&entity) {
82 id
83 } else {
84 let id = self.entity_to_id.len();
85 self.entity_to_id.insert(entity.clone(), id);
86 self.id_to_entity.insert(id, entity);
87 id
88 }
89 }
90
91 fn get_or_create_relation_id(&mut self, relation: String) -> usize {
93 if let Some(&id) = self.relation_to_id.get(&relation) {
94 id
95 } else {
96 let id = self.relation_to_id.len();
97 self.relation_to_id.insert(relation.clone(), id);
98 self.id_to_relation.insert(id, relation);
99 id
100 }
101 }
102
103 pub fn get_entity_id(&self, entity: &str) -> Option<usize> {
105 self.entity_to_id.get(entity).copied()
106 }
107
108 pub fn get_relation_id(&self, relation: &str) -> Option<usize> {
110 self.relation_to_id.get(relation).copied()
111 }
112
113 pub fn get_entity(&self, id: usize) -> Option<&String> {
115 self.id_to_entity.get(&id)
116 }
117
118 pub fn get_relation(&self, id: usize) -> Option<&String> {
120 self.id_to_relation.get(&id)
121 }
122
123 pub fn num_entities(&self) -> usize {
125 self.entity_to_id.len()
126 }
127
128 pub fn num_relations(&self) -> usize {
130 self.relation_to_id.len()
131 }
132
133 pub fn num_triples(&self) -> usize {
135 self.triples.len()
136 }
137
138 pub fn get_entities(&self) -> Vec<String> {
140 self.entity_to_id.keys().cloned().collect()
141 }
142
143 pub fn get_relations(&self) -> Vec<String> {
145 self.relation_to_id.keys().cloned().collect()
146 }
147
148 pub fn has_triple(&self, subject_id: usize, predicate_id: usize, object_id: usize) -> bool {
150 self.positive_triples
151 .contains(&(subject_id, predicate_id, object_id))
152 }
153
154 pub fn generate_negative_samples<R>(
156 &self,
157 num_samples: usize,
158 rng: &mut Random<R>,
159 ) -> Vec<(usize, usize, usize)>
160 where
161 R: scirs2_core::random::Rng,
162 {
163 let mut negative_samples = Vec::new();
164 let num_entities = self.num_entities();
165
166 if self.triples.is_empty() || num_entities == 0 {
172 return negative_samples;
173 }
174
175 let max_attempts = num_samples.saturating_mul(100).max(1000);
180 let mut attempts = 0usize;
181
182 while negative_samples.len() < num_samples && attempts < max_attempts {
183 attempts += 1;
184
185 let idx = rng.random_range(0..self.triples.len());
187 let &(s, p, o) = &self.triples[idx];
188
189 let corrupt_subject = rng.random_bool_with_chance(0.5);
191
192 let negative_triple = if corrupt_subject {
193 let new_subject = rng.random_range(0..num_entities);
194 (new_subject, p, o)
195 } else {
196 let new_object = rng.random_range(0..num_entities);
197 (s, p, new_object)
198 };
199
200 if !self.has_triple(negative_triple.0, negative_triple.1, negative_triple.2) {
202 negative_samples.push(negative_triple);
203 }
204 }
205
206 negative_samples
207 }
208
209 pub fn get_stats(&self, model_type: &str) -> ModelStats {
211 ModelStats {
212 num_entities: self.num_entities(),
213 num_relations: self.num_relations(),
214 num_triples: self.num_triples(),
215 dimensions: self.config.dimensions,
216 is_trained: self.is_trained,
217 model_type: model_type.to_string(),
218 creation_time: self.creation_time,
219 last_training_time: self.last_training_time,
220 }
221 }
222
223 pub fn clear(&mut self) {
225 self.entity_to_id.clear();
226 self.id_to_entity.clear();
227 self.relation_to_id.clear();
228 self.id_to_relation.clear();
229 self.triples.clear();
230 self.positive_triples.clear();
231 self.is_trained = false;
232 self.last_training_time = None;
233 }
234
235 pub fn mark_trained(&mut self) {
237 self.is_trained = true;
238 self.last_training_time = Some(Utc::now());
239 }
240}
241
242#[cfg(test)]
243mod tests {
244 use super::*;
245 use crate::{NamedNode, Triple};
246 use scirs2_core::random::Random;
247
248 #[test]
253 fn regression_negative_sampling_terminates_on_saturated_space() {
254 let mut base = BaseModel::new(ModelConfig::default());
255 let x = NamedNode::new("http://example.org/x").expect("valid iri");
256 let p = NamedNode::new("http://example.org/p").expect("valid iri");
257 base.add_triple(Triple::new(x.clone(), p.clone(), x.clone()))
258 .expect("add triple");
259 assert_eq!(base.num_entities(), 1);
260
261 let mut rng = Random::default();
262 let negatives = base.generate_negative_samples(5, &mut rng);
264 assert!(
265 negatives.len() <= 5,
266 "must not exceed requested count: {}",
267 negatives.len()
268 );
269 }
270
271 #[test]
273 fn regression_negative_sampling_produces_samples_when_possible() {
274 let mut base = BaseModel::new(ModelConfig::default());
275 for i in 0..10 {
276 let s = NamedNode::new(format!("http://example.org/e{i}").as_str()).expect("iri");
277 let p = NamedNode::new("http://example.org/p").expect("iri");
278 let o = NamedNode::new(format!("http://example.org/e{}", i + 1).as_str()).expect("iri");
279 base.add_triple(Triple::new(s, p, o)).expect("add triple");
280 }
281 let mut rng = Random::default();
282 let negatives = base.generate_negative_samples(5, &mut rng);
283 assert_eq!(negatives.len(), 5);
284 }
285}