NOTavaliable commited on
Commit
4be5a03
·
verified ·
1 Parent(s): 331cb33

Upload code/process_sentence2triple.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. code/process_sentence2triple.py +52 -0
code/process_sentence2triple.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ class StochasticRAGTripleExtraction:
2
+ def __init__(self, retriever, generator, temperature=1.0):
3
+ self.retriever = retriever
4
+ self.generator = generator
5
+ self.temperature = temperature
6
+
7
+ def extract_with_sampling(self, text, num_samples=10):
8
+ """
9
+ 通过随机采样多次提取,增强三元组多样性
10
+ """
11
+ all_triples = []
12
+
13
+ for _ in range(num_samples):
14
+ # 随机检索不同的文档
15
+ docs = self.retriever.search(
16
+ text,
17
+ top_k=5,
18
+ temperature=self.temperature # 随机性检索
19
+ )
20
+
21
+ # 随机采样提示模板
22
+ templates = [
23
+ "Extract triples from: {text}",
24
+ "Find all (subject, relation, object) in: {text}",
25
+ "Identify knowledge triples: {text}"
26
+ ]
27
+ template = random.choice(templates)
28
+
29
+ # 生成三元组(带温度参数控制随机性)
30
+ response = self.generator.generate(
31
+ template.format(text=text),
32
+ temperature=self.temperature,
33
+ do_sample=True, # 启用随机采样
34
+ top_p=0.9 # nucleus采样
35
+ )
36
+
37
+ # 解析并收集三元组
38
+ triples = self.parse_triples(response)
39
+ all_triples.extend(triples)
40
+
41
+ # 基于出现频率过滤
42
+ return self.filter_by_frequency(all_triples)
43
+
44
+ def filter_by_frequency(self, triples, min_count=3):
45
+ """
46
+ 保留多次采样中频繁出现的三元组
47
+ """
48
+ triple_counts = Counter(triples)
49
+ return [
50
+ triple for triple, count in triple_counts.items()
51
+ if count >= min_count
52
+ ]