MCPcopy Create free account
hub / github.com/THUDM/AgentBench / KnowledgeGraph

Class KnowledgeGraph

src/server/tasks/knowledgegraph/task.py:26–232  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class KnowledgeGraph(Task):
27
28 def __init__(self,
29 data_file: str,
30 max_rounds: int = 15,
31 one_shot: bool = False,
32 database_file: Optional[str] = None,
33 env_driver: str = 'manual',
34 env_options: Optional[dict] = None,
35 **kwargs):
36 super().__init__(tools=TOOLS, **kwargs)
37 self.logger = logging.getLogger(__name__)
38
39 self.max_rounds = max_rounds
40 self.one_shot = one_shot
41 self.data: List[Tuple[dict, set]] = []
42 self.inputs: List[dict] = []
43 self.targets: List[set] = []
44 with open(data_file, 'r') as f:
45 data_object = json.load(f)
46 for item in data_object:
47 answer = item.pop("answer")
48 gold_answer = set()
49 for a in answer:
50 gold_answer.add(a["answer_argument"])
51 self.data.append((item, gold_answer)) # input and target
52 self.inputs.append(item)
53 self.targets.append(gold_answer)
54
55 self.env_delegation = KnowledgeGraphEnvironmentDelegation(database_file)
56 self.env_controller = create_controller(env_driver, self.env_delegation, **env_options)
57 self.env_controller_background_task = None
58
59 @cache
60 def get_indices(self) -> List[SampleIndex]:
61 return list(range(len(self.data)))
62
63 async def start_sample(self, index: SampleIndex, session: Session) -> TaskSampleExecutionResult:
64 self.env_controller.loop = asyncio.get_running_loop()
65 if not self.env_controller_background_task:
66 self.env_controller_background_task = asyncio.create_task(self.env_controller.background_task())
67 weakref.finalize(self, self.env_controller_background_task.cancel)
68
69 return await super().start_sample(index, session)
70
71 def sync_start_sample(self, index: SampleIndex, session: Session) -> TaskSampleExecutionResult:
72 self.logger.info(f'starting sample {index} with session id {session.id}')
73
74 data_item = self.inputs[index]
75 question = data_item['question']
76 entities = data_item['entities']
77 self.logger.info(f'[session {session.id}] Processing question: {question[:50]}...')
78
79 session_id, _, urls = self.env_controller.sync_start_session(ENV_SUBTYPE)
80 try:
81 sparql_url = urls[ENV_SUBTYPE]
82 sparql_executor = SparqlExecuter(sparql_url)
83 api = API(sparql_executor, session.id)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected