Skip to content

Commit 793e283

Browse files
committed
Added extraction transform
1 parent 21154cb commit 793e283

11 files changed

Lines changed: 212 additions & 85 deletions

src/itp_interface/main/config.py

Lines changed: 1 addition & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,6 @@
11
#!/usr/bin/env python3
22

3-
import sys
4-
5-
root_dir = f"{__file__.split('itp_interface')[0]}"
6-
if root_dir not in sys.path:
7-
sys.path.append(root_dir)
83
import typing
9-
from pydantic import BaseModel
104
from dataclasses import dataclass, field
115
from dataclasses_json import dataclass_json
126
from enum import Enum
@@ -76,7 +70,7 @@ class EvalFile(object):
7670

7771
@dataclass_json
7872
@dataclass
79-
class ExtractFile(BaseModel):
73+
class ExtractFile(object):
8074
path: str
8175
declarations: typing.Union[str, typing.List[str]]
8276

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
name: simple_benchmark_lean_ext
2+
num_files: 1
3+
language: LEAN4
4+
few_shot_data_path_for_retrieval:
5+
few_shot_metadata_filename_for_retrieval:
6+
dfs_data_path_for_retrieval:
7+
dfs_metadata_filename_for_retrieval:
8+
is_extraction_request: true
9+
datasets:
10+
- project: src/data/test/lean4_proj
11+
files:
12+
- path: Lean4Proj/Basic.lean
13+
declarations: "*"
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
defaults:
2+
# - benchmark: simple_benchmark_lean_training_data
3+
# - run_settings: default_lean_data_generation_transforms
4+
# - benchmark: simple_benchmark_1
5+
# - run_settings: default_lean4_data_generation_transforms
6+
- benchmark: simple_benchmark_lean_ext
7+
- run_settings: default_lean4_data_generation_transforms
8+
- env_settings: no_retrieval
9+
- override hydra/job_logging: 'disabled'
10+
11+
run_settings:
12+
output_dir: .log/data_generation/benchmark/simple_benchmark_lean_ext
13+
pool_size: 2

src/itp_interface/main/run_tool.py

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313
import numpy as np
1414
import yaml
1515
import uuid
16-
import threading
1716
from concurrent.futures import ThreadPoolExecutor
1817

1918
# Conditional Ray import
@@ -34,15 +33,14 @@
3433
from itp_interface.tools.coq_local_data_generation_transform import LocalDataGenerationTransform as CoqLocalDataGenerationTransform
3534
from itp_interface.tools.lean_local_data_generation_transform import LocalDataGenerationTransform as LeanLocalDataGenerationTransform
3635
from itp_interface.tools.lean4_local_data_generation_transform import Local4DataGenerationTransform
36+
from itp_interface.tools.lean4_local_data_extraction_transform import Local4DataExtractionTransform
3737
from itp_interface.tools.isabelle_local_data_generation_transform import LocalDataGenerationTransform as IsabelleLocalDataGenerationTransform
3838
from itp_interface.tools.run_data_generation_transforms import RunDataGenerationTransforms
3939
from itp_interface.tools.log_utils import setup_logger
4040
from itp_interface.main.config import EvalFile, ExtractFile, Experiments, EvalRunCheckpointInfo, TransformType, parse_config
41-
from itp_interface.tools.isabelle_executor import IsabelleExecutor
4241
from itp_interface.tools.dynamic_coq_proof_exec import DynamicProofExecutor as DynamicCoqProofExecutor
4342
from itp_interface.tools.dynamic_lean_proof_exec import DynamicProofExecutor as DynamicLeanProofExecutor
4443
from itp_interface.tools.dynamic_lean4_proof_exec import DynamicProofExecutor as DynamicLean4ProofExecutor
45-
from itp_interface.tools.dynamic_isabelle_proof_exec import DynamicProofExecutor as DynamicIsabelleProofExecutor
4644
from itp_interface.tools.coq_executor import get_all_lemmas_in_file as get_all_lemmas_coq
4745
from itp_interface.tools.lean4_sync_executor import get_all_theorems_in_file as get_all_lemmas_lean4, get_fully_qualified_theorem_name as get_fully_qualified_theorem_name_lean4, get_theorem_name_resembling as get_theorem_name_resembling_lean4
4846
from itp_interface.tools.isabelle_executor import get_all_lemmas_in_file as get_all_lemmas_isabelle
@@ -255,11 +253,17 @@ def add_transform(experiment: Experiments, clone_dir: str, resources: list, tran
255253
logger=logger)
256254
os.makedirs(clone_dir, exist_ok=True)
257255
elif experiment.benchmark.language == ProofAction.Language.LEAN4:
258-
transform = Local4DataGenerationTransform(
259-
experiment.run_settings.dep_depth,
260-
max_search_results=experiment.run_settings.max_search_results,
261-
buffer_size=experiment.run_settings.buffer_size,
262-
logger=logger)
256+
if experiment.benchmark.is_extraction_request:
257+
transform = Local4DataExtractionTransform(
258+
experiment.run_settings.dep_depth,
259+
buffer_size=experiment.run_settings.buffer_size,
260+
logger=logger)
261+
else:
262+
transform = Local4DataGenerationTransform(
263+
experiment.run_settings.dep_depth,
264+
max_search_results=experiment.run_settings.max_search_results,
265+
buffer_size=experiment.run_settings.buffer_size,
266+
logger=logger)
263267
clone_dir = None
264268
elif experiment.benchmark.language == ProofAction.Language.COQ:
265269
only_proof_state = experiment.env_settings.retrieval_strategy == ProofEnvReRankStrategy.NO_RE_RANK
@@ -292,7 +296,7 @@ def add_transform(experiment: Experiments, clone_dir: str, resources: list, tran
292296
transforms.append(transform)
293297
else:
294298
raise ValueError(f"Unexpected transform_type: {experiment.run_settings.transform_type}")
295-
pass
299+
return clone_dir
296300

297301
def get_decl_lemmas_to_parse(
298302
experiment: Experiments,
@@ -434,7 +438,7 @@ def run_data_generation_pipeline(experiment: Experiments, log_dir: str, checkpoi
434438
transforms = []
435439
str_time = time.strftime("%Y%m%d-%H%M%S")
436440
clone_dir = os.path.join(experiment.run_settings.output_dir, "clone{}".format(str_time))
437-
add_transform(experiment, clone_dir, resources, transforms, logger)
441+
clone_dir = add_transform(experiment, clone_dir, resources, transforms, logger)
438442
# Find all the lemmas to prove
439443
project_to_theorems = {}
440444
other_args = {}

src/itp_interface/tools/lean4_local_data_extraction_transform.py

Lines changed: 30 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -7,10 +7,9 @@
77
import typing
88
import uuid
99
from itp_interface.tools.simple_lean4_sync_executor import SimpleLean4SyncExecutor
10-
from itp_interface.tools.lean4_context_helper import Lean4ContextHelper
1110
from itp_interface.tools.coq_training_data_generator import GenericTrainingDataGenerationTransform, TrainingDataGenerationType
12-
from itp_interface.tools.training_data_format import MergableCollection, TrainingDataMetadataFormat, TheoremProvingTrainingDataCollection, TheoremProvingTrainingDataFormat
13-
from itp_interface.tools.training_data import TrainingData
11+
from itp_interface.tools.training_data_format import MergableCollection, TrainingDataMetadataFormat, ExtractionDataCollection, TheoremProvingTrainingDataFormat
12+
from itp_interface.tools.training_data import TrainingData, DataLayoutFormat
1413

1514
class Local4DataExtractionTransform(GenericTrainingDataGenerationTransform):
1615
def __init__(self,
@@ -24,70 +23,45 @@ def __init__(self,
2423
self.max_search_results = max_search_results
2524
self.max_parallelism = max_parallelism
2625

27-
def get_meta_object(self) -> MergableCollection:
28-
return TrainingDataMetadataFormat(training_data_buffer_size=self.buffer_size)
26+
def get_meta_object(self) -> TrainingDataMetadataFormat:
27+
return TrainingDataMetadataFormat(
28+
training_data_buffer_size=self.buffer_size,
29+
data_filename_prefix="extraction_data_",
30+
lemma_ref_filename_prefix="extraction_lemma_refs_")
2931

3032
def get_data_collection_object(self) -> MergableCollection:
31-
return TheoremProvingTrainingDataCollection()
33+
return ExtractionDataCollection()
3234

3335
def load_meta_from_file(self, file_path) -> MergableCollection:
3436
return TrainingDataMetadataFormat.load_from_file(file_path)
3537

3638
def load_data_from_file(self, file_path) -> MergableCollection:
37-
return TheoremProvingTrainingDataCollection.load_from_file(file_path, self.logger)
39+
return ExtractionDataCollection.load_from_file(file_path, self.logger)
3840

39-
def __call__(self, training_data: TrainingData, project_id : str, lean_executor: SimpleLean4SyncExecutor, print_coq_executor_callback: typing.Callable[[], SimpleLean4SyncExecutor], theorems: typing.List[str] = None, other_args: dict = {}) -> TrainingData:
40-
print_lean_executor = print_coq_executor_callback()
41-
lean_context_helper = Lean4ContextHelper(print_lean_executor, self.depth, self.logger)
42-
lean_context_helper.__enter__()
41+
def __call__(self,
42+
training_data: TrainingData,
43+
project_id : str,
44+
lean_executor: SimpleLean4SyncExecutor,
45+
print_coq_executor_callback: typing.Callable[[], SimpleLean4SyncExecutor],
46+
theorems: typing.List[str] = None,
47+
other_args: dict = {}) -> TrainingData:
4348
file_namespace = lean_executor.main_file.replace('/', '.')
4449
self.logger.info(f"=========================Processing {file_namespace}=========================")
4550
theorem_id = str(uuid.uuid4())
4651
theorems = set(theorems) if theorems is not None else None
47-
assert len(theorems) == 1, "Only one theorem can be processed at a time"
52+
cnt = 0
4853
with lean_executor:
4954
lean_executor.set_run_exactly()
50-
lean_executor._skip_to_theorem(theorems.pop())
51-
start_goals = lean_context_helper.get_focussed_goals_from_proof_context(lean_executor.proof_context)
52-
ran_next = lean_executor.run_next()
53-
cmd_ran = lean_executor.current_stmt
54-
try:
55-
theorem_id = theorem_id + "/" + lean_executor.get_current_lemma_name()
56-
except:
57-
pass
58-
proof_id = theorem_id
59-
while not lean_executor.execution_complete and ran_next:
60-
if lean_executor.is_in_proof_mode():
61-
end_goals = lean_context_helper.get_focussed_goals_from_proof_context(lean_executor.proof_context)
62-
else:
63-
end_goals = []
64-
if len(start_goals) > 0 and \
65-
(len(start_goals) != len(end_goals) or not all(s_g == e_g for s_g, e_g in zip(start_goals, end_goals))):
66-
tdf = TheoremProvingTrainingDataFormat(
67-
proof_id=proof_id,
68-
all_useful_defns_theorems=[],
69-
start_goals=start_goals,
70-
end_goals=end_goals,
71-
proof_steps=[cmd_ran],
72-
simplified_goals=[],
73-
addition_state_info={},
74-
file_path=lean_executor.main_file,
75-
theorem_name=proof_id,
76-
project_id=project_id)
77-
training_data.merge(tdf)
78-
if lean_executor.is_in_proof_mode():
79-
start_goals = end_goals
80-
ran_next = lean_executor.run_next()
81-
cmd_ran = lean_executor.current_stmt
82-
else:
83-
ran_next = False
84-
training_data.meta.num_theorems += 1
55+
line_infos = lean_executor.extract_all_theorems_and_definitions()
56+
for line_info in line_infos:
57+
if theorems is not None and line_info.name not in theorems:
58+
continue
59+
training_data.merge(line_info)
60+
cnt += 1
61+
training_data.meta.last_proof_id = theorem_id
8562
self.logger.info(f"===============Finished processing {file_namespace}=====================")
86-
self.logger.info(f"Total theorems processed in this transform: 1")
87-
try:
88-
lean_context_helper.__exit__(None, None, None)
89-
except:
90-
pass
63+
self.logger.info(f"Total declarations processed in this transform: {cnt}")
64+
return training_data
9165

9266

9367
if __name__ == "__main__":
@@ -110,13 +84,14 @@ def _print_lean_executor_callback():
11084
search_lean_exec = SimpleLean4SyncExecutor(main_file=file_name, project_root=project_dir)
11185
search_lean_exec.__enter__()
11286
return search_lean_exec
113-
transform = Local4DataGenerationTransform(0, buffer_size=1000)
87+
transform = Local4DataExtractionTransform(0, buffer_size=1000)
11488
training_data = TrainingData(
11589
output_path,
11690
"training_metadata.json",
11791
training_meta=transform.get_meta_object(),
118-
logger=logger)
92+
logger=logger,
93+
layout=DataLayoutFormat.DECLARATION_EXTRACTION)
11994
with SimpleLean4SyncExecutor(project_root=project_dir, main_file=file_name, use_human_readable_proof_context=True, suppress_error_log=True) as coq_exec:
120-
transform(training_data, project_id, coq_exec, _print_lean_executor_callback, theorems=['{"namespace": "Lean4Proj1", "name": "test2"}'])
95+
transform(training_data, project_id, coq_exec, _print_lean_executor_callback, theorems=["test2", "test", "test1"])
12196
save_info = training_data.save()
12297
logger.info(f"Saved training data to {save_info}")

src/itp_interface/tools/run_data_generation_transforms.py

Lines changed: 24 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
import gc
1313
import threading
1414
from concurrent.futures import ThreadPoolExecutor, TimeoutError as FutureTimeoutError
15-
from itp_interface.tools.training_data import TrainingData
15+
from itp_interface.tools.training_data import TrainingData, DataLayoutFormat
1616

1717
# Conditional Ray import
1818
try:
@@ -30,6 +30,7 @@
3030
from itp_interface.tools.coq_local_data_generation_transform import LocalDataGenerationTransform as CoqLocalDataGenerationTransform
3131
from itp_interface.tools.lean_local_data_generation_transform import LocalDataGenerationTransform as LeanLocalDataGenerationTransform
3232
from itp_interface.tools.lean4_local_data_generation_transform import Local4DataGenerationTransform as Lean4LocalDataGenerationTransform
33+
from itp_interface.tools.lean4_local_data_extraction_transform import Local4DataExtractionTransform as Lean4LocalDataExtractionTransform
3334
from itp_interface.tools.isabelle_local_data_generation_transform import LocalDataGenerationTransform as IsabelleLocalDataGenerationTransform
3435
from itp_interface.tools.coq_training_data_generator import GenericTrainingDataGenerationTransform, TrainingDataGenerationType
3536

@@ -144,7 +145,11 @@ def _print_isabelle_callback():
144145
search_isabelle_exec = IsabelleExecutor(project_path, file_path, use_human_readable_proof_context=use_human_readable, suppress_error_log=log_error, port=port)
145146
search_isabelle_exec.__enter__()
146147
return search_isabelle_exec
147-
if isinstance(transform, CoqLocalDataGenerationTransform) or isinstance(transform, LeanLocalDataGenerationTransform) or isinstance(transform, IsabelleLocalDataGenerationTransform) or isinstance(transform, Lean4LocalDataGenerationTransform):
148+
if isinstance(transform, CoqLocalDataGenerationTransform) or \
149+
isinstance(transform, LeanLocalDataGenerationTransform) or \
150+
isinstance(transform, IsabelleLocalDataGenerationTransform) or \
151+
isinstance(transform, Lean4LocalDataGenerationTransform) or \
152+
isinstance(transform, Lean4LocalDataExtractionTransform):
148153
if isinstance(transform, IsabelleLocalDataGenerationTransform) and transform.ray_resource_pool is not None:
149154
# This is a blocking call
150155
port = ray.get(transform.ray_resource_pool.wait_and_acquire.remote(1))[0]
@@ -158,6 +163,8 @@ def _print_isabelle_callback():
158163
exec = IsabelleExecutor(project_path, file_path, use_human_readable_proof_context=use_human_readable, suppress_error_log=log_error, port=port)
159164
elif isinstance(transform, Lean4LocalDataGenerationTransform):
160165
exec = SimpleLean4SyncExecutor(project_path, None, file_path, use_human_readable_proof_context=use_human_readable, suppress_error_log=log_error)
166+
elif isinstance(transform, Lean4LocalDataExtractionTransform):
167+
exec = SimpleLean4SyncExecutor(project_path, None, file_path, use_human_readable_proof_context=use_human_readable, suppress_error_log=log_error)
161168
else:
162169
raise Exception("Unknown transform")
163170
with exec:
@@ -170,6 +177,8 @@ def _print_isabelle_callback():
170177
transform(training_data, project_id, exec, _print_isabelle_callback, theorems, other_args)
171178
elif isinstance(transform, Lean4LocalDataGenerationTransform):
172179
transform(training_data, project_id, exec, _print_lean4_callback, theorems, other_args)
180+
elif isinstance(transform, Lean4LocalDataExtractionTransform):
181+
transform(training_data, project_id, exec, _print_lean4_callback, theorems, other_args)
173182
else:
174183
raise Exception("Unknown transform")
175184
finally:
@@ -195,13 +204,18 @@ def get_training_data_object(transform, output_dir, logger: logging.Logger):
195204
metadata.data_filename_suffix = RunDataGenerationTransforms.get_data_filename_suffix(transform)
196205
metadata.lemma_ref_filename_prefix = RunDataGenerationTransforms.get_lemma_ref_filename_prefix(transform)
197206
metadata.lemma_ref_filename_suffix = RunDataGenerationTransforms.get_lemma_ref_filename_suffix(transform)
207+
if isinstance(transform, Lean4LocalDataExtractionTransform):
208+
layout = DataLayoutFormat.DECLARATION_EXTRACTION
209+
else:
210+
layout = DataLayoutFormat.THEOREM_PROVING
198211
training_data = TrainingData(
199212
output_dir,
200213
RunDataGenerationTransforms.get_meta_file_name(transform),
201214
metadata,
202215
transform.max_parallelism,
203216
remove_from_store_after_loading=True,
204-
logger=logger)
217+
logger=logger,
218+
layout=layout)
205219
return training_data
206220

207221
@staticmethod
@@ -296,7 +310,7 @@ def run_local_transform(self, pool_size: int , transform: typing.Union[CoqLocalD
296310
os.makedirs(temp_file_dir, exist_ok=True)
297311
log_file = os.path.join(self.logging_dir, f"{relative_file_path}.log")
298312
theorems = projects[project][file_path]
299-
if isinstance(transform, Lean4LocalDataGenerationTransform):
313+
if isinstance(transform, Lean4LocalDataGenerationTransform) or isinstance(transform, Lean4LocalDataExtractionTransform):
300314
# For every theorem we need to create a separate job
301315
for _idx, theorem in enumerate(theorems):
302316
log_file = os.path.join(self.logging_dir, f"{relative_file_path}-{_idx}.log")
@@ -317,13 +331,18 @@ def run_local_transform(self, pool_size: int , transform: typing.Union[CoqLocalD
317331
final_training_meta.data_filename_suffix = RunDataGenerationTransforms.get_data_filename_suffix(transform)
318332
final_training_meta.lemma_ref_filename_prefix = RunDataGenerationTransforms.get_lemma_ref_filename_prefix(transform)
319333
final_training_meta.lemma_ref_filename_suffix = RunDataGenerationTransforms.get_lemma_ref_filename_suffix(transform)
334+
if isinstance(transform, Lean4LocalDataExtractionTransform):
335+
layout = DataLayoutFormat.DECLARATION_EXTRACTION
336+
else:
337+
layout = DataLayoutFormat.THEOREM_PROVING
320338
final_training_data = TrainingData(
321339
new_output_dir,
322340
RunDataGenerationTransforms.get_meta_file_name(transform),
323341
final_training_meta,
324342
transform.max_parallelism,
325343
remove_from_store_after_loading=True,
326-
logger=self.logger)
344+
logger=self.logger,
345+
layout=layout)
327346
last_job_idx = 0
328347
tds = [None]*len(job_spec)
329348
num_theorems = 0

0 commit comments

Comments
 (0)