(self,
data_file: str,
max_rounds: int = 15,
one_shot: bool = False,
database_file: Optional[str] = None,
env_driver: str = 'manual',
env_options: Optional[dict] = None,
**kwargs)
| 26 | class 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]: |
nothing calls this directly
no test coverage detected