1use crate::models::serialization::{BaseModelSnapshot, MatrixF64, Tensor3F64};
9use crate::models::{common::*, BaseModel};
10use crate::{EmbeddingModel, ModelConfig, ModelStats, TrainingStats, Triple, Vector};
11use anyhow::{anyhow, Result};
12use async_trait::async_trait;
13use scirs2_core::ndarray_ext::{Array2, Array3};
14use scirs2_core::random::{Random, Rng, RngExt, SliceRandom};
15use serde::{Deserialize, Serialize};
16use std::fs::File;
17use std::io::{BufReader, BufWriter};
18use std::path::Path;
19use std::time::Instant;
20use tracing::{debug, info};
21use uuid::Uuid;
22
23#[derive(Debug, Serialize, Deserialize)]
25struct TuckERSerializable {
26 base: BaseModelSnapshot,
27 entity_embeddings: MatrixF64,
28 relation_embeddings: MatrixF64,
29 core_tensor: Tensor3F64,
30 embeddings_initialized: bool,
31 entity_dim: usize,
32 relation_dim: usize,
33 core_dims: (usize, usize, usize),
34 dropout_rate: f64,
35 batch_norm: bool,
36}
37
38#[derive(Debug)]
40pub struct TuckER {
41 base: BaseModel,
43 entity_embeddings: Array2<f64>,
45 relation_embeddings: Array2<f64>,
47 core_tensor: Array3<f64>,
49 embeddings_initialized: bool,
51 entity_dim: usize,
53 relation_dim: usize,
55 core_dims: (usize, usize, usize),
57 dropout_rate: f64,
59 batch_norm: bool,
61}
62
63impl TuckER {
64 pub fn new(config: ModelConfig) -> Self {
66 let base = BaseModel::new(config.clone());
67
68 let entity_dim = config
70 .model_params
71 .get("entity_dim")
72 .map(|&v| v as usize)
73 .unwrap_or(config.dimensions);
74 let relation_dim = config
75 .model_params
76 .get("relation_dim")
77 .map(|&v| v as usize)
78 .unwrap_or(config.dimensions);
79 let core_dim1 = config
80 .model_params
81 .get("core_dim1")
82 .map(|&v| v as usize)
83 .unwrap_or(config.dimensions);
84 let core_dim2 = config
85 .model_params
86 .get("core_dim2")
87 .map(|&v| v as usize)
88 .unwrap_or(config.dimensions);
89 let core_dim3 = config
90 .model_params
91 .get("core_dim3")
92 .map(|&v| v as usize)
93 .unwrap_or(config.dimensions);
94 let dropout_rate = config
95 .model_params
96 .get("dropout_rate")
97 .copied()
98 .unwrap_or(0.3);
99 let batch_norm = config
100 .model_params
101 .get("batch_norm")
102 .map(|&v| v > 0.0)
103 .unwrap_or(true);
104
105 Self {
106 base,
107 entity_embeddings: Array2::zeros((0, entity_dim)),
108 relation_embeddings: Array2::zeros((0, relation_dim)),
109 core_tensor: Array3::zeros((core_dim1, core_dim2, core_dim3)),
110 embeddings_initialized: false,
111 entity_dim,
112 relation_dim,
113 core_dims: (core_dim1, core_dim2, core_dim3),
114 dropout_rate,
115 batch_norm,
116 }
117 }
118
119 fn initialize_embeddings(&mut self) {
121 if self.embeddings_initialized {
122 return;
123 }
124
125 let num_entities = self.base.num_entities();
126 let num_relations = self.base.num_relations();
127
128 if num_entities == 0 || num_relations == 0 {
129 return;
130 }
131
132 let mut rng = Random::seed(self.base.config.seed.unwrap_or_else(|| {
133 use std::time::{SystemTime, UNIX_EPOCH};
134 SystemTime::now()
135 .duration_since(UNIX_EPOCH)
136 .expect("SystemTime should be after UNIX_EPOCH")
137 .as_secs()
138 }));
139
140 self.entity_embeddings = xavier_init(
142 (num_entities, self.entity_dim),
143 self.entity_dim,
144 self.entity_dim,
145 &mut rng,
146 );
147
148 self.relation_embeddings = xavier_init(
150 (num_relations, self.relation_dim),
151 self.relation_dim,
152 self.relation_dim,
153 &mut rng,
154 );
155
156 let total_elements = self.core_dims.0 * self.core_dims.1 * self.core_dims.2;
158 let std_dev = (2.0 / total_elements as f64).sqrt();
159
160 for elem in self.core_tensor.iter_mut() {
161 *elem = rng.random_range(-std_dev..std_dev);
162 }
163
164 normalize_embeddings(&mut self.entity_embeddings);
166 normalize_embeddings(&mut self.relation_embeddings);
167
168 self.embeddings_initialized = true;
169 debug!(
170 "Initialized TuckER embeddings: {} entities ({}D), {} relations ({}D), core tensor {:?}",
171 num_entities, self.entity_dim, num_relations, self.relation_dim, self.core_dims
172 );
173 }
174
175 fn score_triple_ids(
177 &self,
178 subject_id: usize,
179 predicate_id: usize,
180 object_id: usize,
181 ) -> Result<f64> {
182 if !self.embeddings_initialized {
183 return Err(anyhow!("Model not trained"));
184 }
185
186 let h = self.entity_embeddings.row(subject_id);
187 let r = self.relation_embeddings.row(predicate_id);
188 let t = self.entity_embeddings.row(object_id);
189
190 let mut score = 0.0;
193
194 for i in 0..self.core_dims.0.min(h.len()) {
195 for j in 0..self.core_dims.1.min(r.len()) {
196 for k in 0..self.core_dims.2.min(t.len()) {
197 score += h[i] * r[j] * t[k] * self.core_tensor[(i, j, k)];
198 }
199 }
200 }
201
202 Ok(score)
203 }
204
205 fn compute_gradients(
207 &self,
208 pos_triple: (usize, usize, usize),
209 neg_triple: (usize, usize, usize),
210 _learning_rate: f64,
211 ) -> Result<(Array2<f64>, Array2<f64>, Array3<f64>)> {
212 let (pos_s, pos_p, pos_o) = pos_triple;
213 let (neg_s, neg_p, neg_o) = neg_triple;
214
215 let mut entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
216 let mut relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
217 let mut core_grads = Array3::zeros(self.core_tensor.raw_dim());
218
219 let pos_score = self.score_triple_ids(pos_s, pos_p, pos_o)?;
221 let neg_score = self.score_triple_ids(neg_s, neg_p, neg_o)?;
222
223 let pos_sigmoid = 1.0 / (1.0 + (-pos_score).exp());
225 let neg_sigmoid = 1.0 / (1.0 + (-neg_score).exp());
226
227 let pos_grad = pos_sigmoid - 1.0;
228 let neg_grad = neg_sigmoid;
229
230 self.compute_triple_gradients(
232 pos_triple,
233 pos_grad,
234 &mut entity_grads,
235 &mut relation_grads,
236 &mut core_grads,
237 );
238
239 self.compute_triple_gradients(
241 neg_triple,
242 neg_grad,
243 &mut entity_grads,
244 &mut relation_grads,
245 &mut core_grads,
246 );
247
248 Ok((entity_grads, relation_grads, core_grads))
249 }
250
251 fn compute_triple_gradients(
253 &self,
254 triple: (usize, usize, usize),
255 loss_grad: f64,
256 entity_grads: &mut Array2<f64>,
257 relation_grads: &mut Array2<f64>,
258 core_grads: &mut Array3<f64>,
259 ) {
260 let (s, p, o) = triple;
261
262 let h = self.entity_embeddings.row(s);
263 let r = self.relation_embeddings.row(p);
264 let t = self.entity_embeddings.row(o);
265
266 for i in 0..self.core_dims.0.min(h.len()) {
268 let mut h_grad = 0.0;
269 for j in 0..self.core_dims.1.min(r.len()) {
270 for k in 0..self.core_dims.2.min(t.len()) {
271 h_grad += r[j] * t[k] * self.core_tensor[(i, j, k)];
272 }
273 }
274 entity_grads[[s, i]] += loss_grad * h_grad;
275 }
276
277 for k in 0..self.core_dims.2.min(t.len()) {
278 let mut t_grad = 0.0;
279 for i in 0..self.core_dims.0.min(h.len()) {
280 for j in 0..self.core_dims.1.min(r.len()) {
281 t_grad += h[i] * r[j] * self.core_tensor[(i, j, k)];
282 }
283 }
284 entity_grads[[o, k]] += loss_grad * t_grad;
285 }
286
287 for j in 0..self.core_dims.1.min(r.len()) {
289 let mut r_grad = 0.0;
290 for i in 0..self.core_dims.0.min(h.len()) {
291 for k in 0..self.core_dims.2.min(t.len()) {
292 r_grad += h[i] * t[k] * self.core_tensor[(i, j, k)];
293 }
294 }
295 relation_grads[[p, j]] += loss_grad * r_grad;
296 }
297
298 for i in 0..self.core_dims.0.min(h.len()) {
300 for j in 0..self.core_dims.1.min(r.len()) {
301 for k in 0..self.core_dims.2.min(t.len()) {
302 core_grads[[i, j, k]] += loss_grad * h[i] * r[j] * t[k];
303 }
304 }
305 }
306 }
307
308 async fn train_epoch(&mut self, learning_rate: f64) -> Result<f64> {
310 let mut rng = Random::seed(self.base.config.seed.unwrap_or_else(|| {
311 use std::time::{SystemTime, UNIX_EPOCH};
312 SystemTime::now()
313 .duration_since(UNIX_EPOCH)
314 .expect("SystemTime should be after UNIX_EPOCH")
315 .as_secs()
316 }));
317
318 let mut total_loss = 0.0;
319 let num_batches = (self.base.triples.len() + self.base.config.batch_size - 1)
320 / self.base.config.batch_size;
321
322 let mut shuffled_triples = self.base.triples.clone();
324 shuffled_triples.shuffle(&mut rng);
325
326 for batch_triples in shuffled_triples.chunks(self.base.config.batch_size) {
327 let mut batch_entity_grads = Array2::zeros(self.entity_embeddings.raw_dim());
328 let mut batch_relation_grads = Array2::zeros(self.relation_embeddings.raw_dim());
329 let mut batch_core_grads = Array3::zeros(self.core_tensor.raw_dim());
330 let mut batch_loss = 0.0;
331
332 for &pos_triple in batch_triples {
333 let neg_samples = self
335 .base
336 .generate_negative_samples(self.base.config.negative_samples, &mut rng);
337
338 for neg_triple in neg_samples {
339 let pos_score =
341 self.score_triple_ids(pos_triple.0, pos_triple.1, pos_triple.2)?;
342 let neg_score =
343 self.score_triple_ids(neg_triple.0, neg_triple.1, neg_triple.2)?;
344
345 let pos_loss = -(1.0 / (1.0 + (-pos_score).exp())).ln();
347 let neg_loss = -(1.0 / (1.0 + neg_score.exp())).ln();
348 let loss = pos_loss + neg_loss;
349 batch_loss += loss;
350
351 let (entity_grads, relation_grads, core_grads) =
353 self.compute_gradients(pos_triple, neg_triple, learning_rate)?;
354
355 batch_entity_grads += &entity_grads;
356 batch_relation_grads += &relation_grads;
357 batch_core_grads += &core_grads;
358 }
359 }
360
361 if batch_loss > 0.0 {
363 gradient_update(
364 &mut self.entity_embeddings,
365 &batch_entity_grads,
366 learning_rate,
367 self.base.config.l2_reg,
368 );
369
370 gradient_update(
371 &mut self.relation_embeddings,
372 &batch_relation_grads,
373 learning_rate,
374 self.base.config.l2_reg,
375 );
376
377 for ((_i, _j, _k), value) in self.core_tensor.indexed_iter_mut() {
379 let reg_term = self.base.config.l2_reg * *value;
382 *value -= learning_rate * reg_term;
383 }
384
385 if self.dropout_rate > 0.0 {
387 apply_dropout(&mut self.entity_embeddings, self.dropout_rate, &mut rng);
388 apply_dropout(&mut self.relation_embeddings, self.dropout_rate, &mut rng);
389 }
390
391 normalize_embeddings(&mut self.entity_embeddings);
393 normalize_embeddings(&mut self.relation_embeddings);
394 }
395
396 total_loss += batch_loss;
397 }
398
399 Ok(total_loss / num_batches as f64)
400 }
401}
402
403#[async_trait]
404impl EmbeddingModel for TuckER {
405 fn config(&self) -> &ModelConfig {
406 &self.base.config
407 }
408
409 fn model_id(&self) -> &Uuid {
410 &self.base.model_id
411 }
412
413 fn model_type(&self) -> &'static str {
414 "TuckER"
415 }
416
417 fn add_triple(&mut self, triple: Triple) -> Result<()> {
418 self.base.add_triple(triple)
419 }
420
421 async fn train(&mut self, epochs: Option<usize>) -> Result<TrainingStats> {
422 let start_time = Instant::now();
423 let max_epochs = epochs.unwrap_or(self.base.config.max_epochs);
424
425 self.initialize_embeddings();
427
428 if !self.embeddings_initialized {
429 return Err(anyhow!("No training data available"));
430 }
431
432 let mut loss_history = Vec::new();
433 let learning_rate = self.base.config.learning_rate;
434
435 info!("Starting TuckER training for {} epochs", max_epochs);
436
437 for epoch in 0..max_epochs {
438 let epoch_loss = self.train_epoch(learning_rate).await?;
439 loss_history.push(epoch_loss);
440
441 if epoch % 100 == 0 {
442 debug!("Epoch {}: loss = {:.6}", epoch, epoch_loss);
443 }
444
445 if epoch > 10 && epoch_loss < 1e-6 {
447 info!("Converged at epoch {} with loss {:.6}", epoch, epoch_loss);
448 break;
449 }
450 }
451
452 self.base.mark_trained();
453 let training_time = start_time.elapsed().as_secs_f64();
454
455 Ok(TrainingStats {
456 epochs_completed: loss_history.len(),
457 final_loss: loss_history.last().copied().unwrap_or(0.0),
458 training_time_seconds: training_time,
459 convergence_achieved: loss_history.last().copied().unwrap_or(f64::INFINITY) < 1e-6,
460 loss_history,
461 })
462 }
463
464 fn get_entity_embedding(&self, entity: &str) -> Result<Vector> {
465 if !self.embeddings_initialized {
466 return Err(anyhow!("Model not trained"));
467 }
468
469 let entity_id = self
470 .base
471 .get_entity_id(entity)
472 .ok_or_else(|| anyhow!("Entity not found: {}", entity))?;
473
474 let embedding = self.entity_embeddings.row(entity_id).to_owned();
475 Ok(ndarray_to_vector(&embedding))
476 }
477
478 fn get_relation_embedding(&self, relation: &str) -> Result<Vector> {
479 if !self.embeddings_initialized {
480 return Err(anyhow!("Model not trained"));
481 }
482
483 let relation_id = self
484 .base
485 .get_relation_id(relation)
486 .ok_or_else(|| anyhow!("Relation not found: {}", relation))?;
487
488 let embedding = self.relation_embeddings.row(relation_id).to_owned();
489 Ok(ndarray_to_vector(&embedding))
490 }
491
492 fn score_triple(&self, subject: &str, predicate: &str, object: &str) -> Result<f64> {
493 let subject_id = self
494 .base
495 .get_entity_id(subject)
496 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
497 let predicate_id = self
498 .base
499 .get_relation_id(predicate)
500 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
501 let object_id = self
502 .base
503 .get_entity_id(object)
504 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
505
506 self.score_triple_ids(subject_id, predicate_id, object_id)
507 }
508
509 fn predict_objects(
510 &self,
511 subject: &str,
512 predicate: &str,
513 k: usize,
514 ) -> Result<Vec<(String, f64)>> {
515 if !self.embeddings_initialized {
516 return Err(anyhow!("Model not trained"));
517 }
518
519 let subject_id = self
520 .base
521 .get_entity_id(subject)
522 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
523 let predicate_id = self
524 .base
525 .get_relation_id(predicate)
526 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
527
528 let mut scores = Vec::new();
529
530 for object_id in 0..self.base.num_entities() {
531 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
532 let object_name = self
533 .base
534 .get_entity(object_id)
535 .expect("entity should exist for valid id")
536 .clone();
537 scores.push((object_name, score));
538 }
539
540 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
541 scores.truncate(k);
542
543 Ok(scores)
544 }
545
546 fn predict_subjects(
547 &self,
548 predicate: &str,
549 object: &str,
550 k: usize,
551 ) -> Result<Vec<(String, f64)>> {
552 if !self.embeddings_initialized {
553 return Err(anyhow!("Model not trained"));
554 }
555
556 let predicate_id = self
557 .base
558 .get_relation_id(predicate)
559 .ok_or_else(|| anyhow!("Predicate not found: {}", predicate))?;
560 let object_id = self
561 .base
562 .get_entity_id(object)
563 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
564
565 let mut scores = Vec::new();
566
567 for subject_id in 0..self.base.num_entities() {
568 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
569 let subject_name = self
570 .base
571 .get_entity(subject_id)
572 .expect("entity should exist for valid id")
573 .clone();
574 scores.push((subject_name, score));
575 }
576
577 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
578 scores.truncate(k);
579
580 Ok(scores)
581 }
582
583 fn predict_relations(
584 &self,
585 subject: &str,
586 object: &str,
587 k: usize,
588 ) -> Result<Vec<(String, f64)>> {
589 if !self.embeddings_initialized {
590 return Err(anyhow!("Model not trained"));
591 }
592
593 let subject_id = self
594 .base
595 .get_entity_id(subject)
596 .ok_or_else(|| anyhow!("Subject not found: {}", subject))?;
597 let object_id = self
598 .base
599 .get_entity_id(object)
600 .ok_or_else(|| anyhow!("Object not found: {}", object))?;
601
602 let mut scores = Vec::new();
603
604 for predicate_id in 0..self.base.num_relations() {
605 let score = self.score_triple_ids(subject_id, predicate_id, object_id)?;
606 let predicate_name = self
607 .base
608 .get_relation(predicate_id)
609 .expect("relation should exist for valid id")
610 .clone();
611 scores.push((predicate_name, score));
612 }
613
614 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
615 scores.truncate(k);
616
617 Ok(scores)
618 }
619
620 fn get_entities(&self) -> Vec<String> {
621 self.base.get_entities()
622 }
623
624 fn get_relations(&self) -> Vec<String> {
625 self.base.get_relations()
626 }
627
628 fn get_stats(&self) -> ModelStats {
629 self.base.get_stats("TuckER")
630 }
631
632 fn save(&self, path: &str) -> Result<()> {
633 info!("Saving TuckER model to {}", path);
634
635 let serializable = TuckERSerializable {
636 base: BaseModelSnapshot::capture(&self.base),
637 entity_embeddings: MatrixF64::from_array(&self.entity_embeddings),
638 relation_embeddings: MatrixF64::from_array(&self.relation_embeddings),
639 core_tensor: Tensor3F64::from_array(&self.core_tensor),
640 embeddings_initialized: self.embeddings_initialized,
641 entity_dim: self.entity_dim,
642 relation_dim: self.relation_dim,
643 core_dims: self.core_dims,
644 dropout_rate: self.dropout_rate,
645 batch_norm: self.batch_norm,
646 };
647
648 let file = File::create(path)
649 .map_err(|e| anyhow!("Failed to create model file {}: {}", path, e))?;
650 let writer = BufWriter::new(file);
651 oxicode::serde::encode_into_std_write(&serializable, writer, oxicode::config::standard())
652 .map_err(|e| anyhow!("Failed to serialize TuckER model: {}", e))?;
653
654 info!("TuckER model saved successfully");
655 Ok(())
656 }
657
658 fn load(&mut self, path: &str) -> Result<()> {
659 info!("Loading TuckER model from {}", path);
660
661 if !Path::new(path).exists() {
662 return Err(anyhow!("Model file not found: {}", path));
663 }
664
665 let file =
666 File::open(path).map_err(|e| anyhow!("Failed to open model file {}: {}", path, e))?;
667 let reader = BufReader::new(file);
668 let (serializable, _): (TuckERSerializable, _) =
669 oxicode::serde::decode_from_std_read(reader, oxicode::config::standard())
670 .map_err(|e| anyhow!("Failed to deserialize TuckER model: {}", e))?;
671
672 self.entity_embeddings = serializable.entity_embeddings.to_array()?;
673 self.relation_embeddings = serializable.relation_embeddings.to_array()?;
674 self.core_tensor = serializable.core_tensor.to_array()?;
675 self.embeddings_initialized = serializable.embeddings_initialized;
676 self.entity_dim = serializable.entity_dim;
677 self.relation_dim = serializable.relation_dim;
678 self.core_dims = serializable.core_dims;
679 self.dropout_rate = serializable.dropout_rate;
680 self.batch_norm = serializable.batch_norm;
681 serializable.base.restore_into(&mut self.base);
682
683 info!("TuckER model loaded successfully");
684 Ok(())
685 }
686
687 fn clear(&mut self) {
688 self.base.clear();
689 self.entity_embeddings = Array2::zeros((0, self.entity_dim));
690 self.relation_embeddings = Array2::zeros((0, self.relation_dim));
691 self.core_tensor = Array3::zeros(self.core_dims);
692 self.embeddings_initialized = false;
693 }
694
695 fn is_trained(&self) -> bool {
696 self.base.is_trained
697 }
698
699 async fn encode(&self, _texts: &[String]) -> Result<Vec<Vec<f32>>> {
700 Err(anyhow!(
701 "Knowledge graph embedding model does not support text encoding"
702 ))
703 }
704}
705
706fn apply_dropout<R: Rng>(embeddings: &mut Array2<f64>, dropout_rate: f64, rng: &mut Random<R>) {
708 for elem in embeddings.iter_mut() {
709 if rng.random::<f64>() < dropout_rate {
710 *elem = 0.0;
711 } else {
712 *elem /= 1.0 - dropout_rate;
713 }
714 }
715}
716
717#[cfg(test)]
718mod tests {
719 use super::*;
720 use crate::NamedNode;
721
722 #[tokio::test]
723 #[cfg_attr(debug_assertions, ignore = "Training tests require release builds")]
724 async fn test_tucker_basic() -> Result<()> {
725 let mut config = ModelConfig::default()
726 .with_dimensions(50)
727 .with_max_epochs(10)
728 .with_seed(42);
729
730 config.model_params.insert("entity_dim".to_string(), 50.0);
732 config.model_params.insert("relation_dim".to_string(), 50.0);
733 config.model_params.insert("core_dim1".to_string(), 50.0);
734 config.model_params.insert("core_dim2".to_string(), 50.0);
735 config.model_params.insert("core_dim3".to_string(), 50.0);
736 config.model_params.insert("dropout_rate".to_string(), 0.1);
737
738 let mut model = TuckER::new(config);
739
740 let alice = NamedNode::new("http://example.org/alice")?;
742 let knows = NamedNode::new("http://example.org/knows")?;
743 let bob = NamedNode::new("http://example.org/bob")?;
744
745 model.add_triple(Triple::new(alice.clone(), knows.clone(), bob.clone()))?;
746 model.add_triple(Triple::new(bob.clone(), knows.clone(), alice.clone()))?;
747
748 let stats = model.train(Some(5)).await?;
750 assert!(stats.epochs_completed > 0);
751
752 let alice_emb = model.get_entity_embedding("http://example.org/alice")?;
754 assert_eq!(alice_emb.dimensions, 50);
755
756 let score = model.score_triple(
758 "http://example.org/alice",
759 "http://example.org/knows",
760 "http://example.org/bob",
761 )?;
762
763 assert!(score.is_finite());
765
766 Ok(())
767 }
768
769 #[test]
770 fn test_tucker_creation() {
771 let config = ModelConfig::default();
772 let tucker = TuckER::new(config);
773 assert!(!tucker.embeddings_initialized);
774 assert_eq!(tucker.model_type(), "TuckER");
775 }
776}