Skip to content

Commit 8fa5bc8

Browse files
fix: optimize import modules
1 parent 661211e commit 8fa5bc8

25 files changed

Lines changed: 426 additions & 187 deletions

graphgen/bases/base_filter.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,8 @@
11
from abc import ABC, abstractmethod
2-
from typing import Any, Union
2+
from typing import TYPE_CHECKING, Any, Union
33

4-
import numpy as np
4+
if TYPE_CHECKING:
5+
import numpy as np
56

67

78
class BaseFilter(ABC):
@@ -15,7 +16,7 @@ def filter(self, data: Any) -> bool:
1516

1617
class BaseValueFilter(BaseFilter, ABC):
1718
@abstractmethod
18-
def filter(self, data: Union[int, float, np.number]) -> bool:
19+
def filter(self, data: Union[int, float, "np.number"]) -> bool:
1920
"""
2021
Filter the numeric value and return True if it passes the filter, False otherwise.
2122
"""

graphgen/bases/base_operator.py

Lines changed: 17 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,18 @@
1+
from __future__ import annotations
2+
13
import inspect
24
import os
35
from abc import ABC, abstractmethod
4-
from typing import Iterable, Tuple, Union
6+
from typing import TYPE_CHECKING, Iterable, Tuple, Union
57

6-
import numpy as np
7-
import pandas as pd
8-
import ray
8+
if TYPE_CHECKING:
9+
import numpy as np
10+
import pandas as pd
911

1012

1113
def convert_to_serializable(obj):
14+
import numpy as np
15+
1216
if isinstance(obj, np.ndarray):
1317
return obj.tolist()
1418
if isinstance(obj, np.generic):
@@ -40,6 +44,8 @@ def __init__(
4044
)
4145

4246
try:
47+
import ray
48+
4349
ctx = ray.get_runtime_context()
4450
worker_id = ctx.get_actor_id() or ctx.get_worker_id()
4551
worker_id_short = worker_id[-6:] if worker_id else "driver"
@@ -62,9 +68,11 @@ def __init__(
6268
)
6369

6470
def __call__(
65-
self, batch: pd.DataFrame
66-
) -> Union[pd.DataFrame, Iterable[pd.DataFrame]]:
71+
self, batch: "pd.DataFrame"
72+
) -> Union["pd.DataFrame", Iterable["pd.DataFrame"]]:
6773
# lazy import to avoid circular import
74+
import pandas as pd
75+
6876
from graphgen.utils import CURRENT_LOGGER_VAR
6977

7078
logger_token = CURRENT_LOGGER_VAR.set(self.logger)
@@ -106,14 +114,16 @@ def get_trace_id(self, content: dict) -> str:
106114

107115
return compute_dict_hash(content, prefix=f"{self.op_name}-")
108116

109-
def split(self, batch: pd.DataFrame) -> tuple[pd.DataFrame, pd.DataFrame]:
117+
def split(self, batch: "pd.DataFrame") -> tuple["pd.DataFrame", "pd.DataFrame"]:
110118
"""
111119
Split the input batch into to_process & processed based on _meta data in KV_storage
112120
:param batch
113121
:return:
114122
to_process: DataFrame of documents to be chunked
115123
recovered: Result DataFrame of already chunked documents
116124
"""
125+
import pandas as pd
126+
117127
meta_forward = self.get_meta_forward()
118128
meta_ids = set(meta_forward.keys())
119129
mask = batch["_trace_id"].isin(meta_ids)

graphgen/bases/base_reader.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,14 @@
1+
from __future__ import annotations
2+
13
import os
24
from abc import ABC, abstractmethod
3-
from typing import Any, Dict, List, Union
5+
from typing import TYPE_CHECKING, Any, Dict, List, Union
46

5-
import pandas as pd
67
import requests
7-
from ray.data import Dataset
8+
9+
if TYPE_CHECKING:
10+
import pandas as pd
11+
from ray.data import Dataset
812

913

1014
class BaseReader(ABC):
@@ -51,6 +55,7 @@ def _validate_batch(self, batch: pd.DataFrame) -> pd.DataFrame:
5155
"""
5256
Validate data format.
5357
"""
58+
5459
if "type" not in batch.columns:
5560
raise ValueError(f"Missing 'type' column. Found: {list(batch.columns)}")
5661

graphgen/common/init_llm.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,12 @@
11
import os
2-
from typing import Any, Dict, Optional
3-
4-
import ray
2+
from typing import TYPE_CHECKING, Any, Dict, Optional
53

64
from graphgen.bases import BaseLLMWrapper
75
from graphgen.models import Tokenizer
86

7+
if TYPE_CHECKING:
8+
import ray
9+
910

1011
class LLMServiceActor:
1112
"""
@@ -73,7 +74,7 @@ class LLMServiceProxy(BaseLLMWrapper):
7374
A proxy class to interact with the LLMServiceActor for distributed LLM operations.
7475
"""
7576

76-
def __init__(self, actor_handle: ray.actor.ActorHandle):
77+
def __init__(self, actor_handle: "ray.actor.ActorHandle"):
7778
super().__init__()
7879
self.actor_handle = actor_handle
7980
self._create_local_tokenizer()
@@ -120,6 +121,8 @@ class LLMFactory:
120121
def create_llm(
121122
model_type: str, backend: str, config: Dict[str, Any]
122123
) -> BaseLLMWrapper:
124+
import ray
125+
123126
if not config:
124127
raise ValueError(
125128
f"No configuration provided for LLM {model_type} with backend {backend}."

graphgen/common/init_storage.py

Lines changed: 76 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
1-
from typing import Any, Dict, List, Set, Union
2-
3-
import ray
1+
from typing import TYPE_CHECKING, Any, Dict, List, Set, Union
42

53
from graphgen.bases.base_storage import BaseGraphStorage, BaseKVStorage
64

5+
if TYPE_CHECKING:
6+
import ray
7+
78

89
class KVStorageActor:
910
def __init__(self, backend: str, working_dir: str, namespace: str):
@@ -146,124 +147,192 @@ def ready(self) -> bool:
146147

147148

148149
class RemoteKVStorageProxy(BaseKVStorage):
149-
def __init__(self, actor_handle: ray.actor.ActorHandle):
150+
def __init__(self, actor_handle: "ray.actor.ActorHandle"):
150151
super().__init__()
151152
self.actor = actor_handle
152153

153154
def data(self) -> Dict[str, Any]:
155+
import ray
156+
154157
return ray.get(self.actor.data.remote())
155158

156159
def all_keys(self) -> list[str]:
160+
import ray
161+
157162
return ray.get(self.actor.all_keys.remote())
158163

159164
def index_done_callback(self):
165+
import ray
166+
160167
return ray.get(self.actor.index_done_callback.remote())
161168

162169
def get_by_id(self, id: str) -> Union[Any, None]:
170+
import ray
171+
163172
return ray.get(self.actor.get_by_id.remote(id))
164173

165174
def get_by_ids(self, ids: list[str], fields=None) -> list[Any]:
175+
import ray
176+
166177
return ray.get(self.actor.get_by_ids.remote(ids, fields))
167178

168179
def get_all(self) -> Dict[str, Any]:
180+
import ray
181+
169182
return ray.get(self.actor.get_all.remote())
170183

171184
def filter_keys(self, data: list[str]) -> set[str]:
185+
import ray
186+
172187
return ray.get(self.actor.filter_keys.remote(data))
173188

174189
def upsert(self, data: Dict[str, Any]):
190+
import ray
191+
175192
return ray.get(self.actor.upsert.remote(data))
176193

177194
def update(self, data: Dict[str, Any]):
195+
import ray
196+
178197
return ray.get(self.actor.update.remote(data))
179198

180199
def delete(self, ids: list[str]):
200+
import ray
201+
181202
return ray.get(self.actor.delete.remote(ids))
182203

183204
def drop(self):
205+
import ray
206+
184207
return ray.get(self.actor.drop.remote())
185208

186209
def reload(self):
210+
import ray
211+
187212
return ray.get(self.actor.reload.remote())
188213

189214

190215
class RemoteGraphStorageProxy(BaseGraphStorage):
191-
def __init__(self, actor_handle: ray.actor.ActorHandle):
216+
def __init__(self, actor_handle: "ray.actor.ActorHandle"):
192217
super().__init__()
193218
self.actor = actor_handle
194219

195220
def index_done_callback(self):
221+
import ray
222+
196223
return ray.get(self.actor.index_done_callback.remote())
197224

198225
def is_directed(self) -> bool:
226+
import ray
227+
199228
return ray.get(self.actor.is_directed.remote())
200229

201230
def get_all_node_degrees(self) -> Dict[str, int]:
231+
import ray
232+
202233
return ray.get(self.actor.get_all_node_degrees.remote())
203234

204235
def get_node_count(self) -> int:
236+
import ray
237+
205238
return ray.get(self.actor.get_node_count.remote())
206239

207240
def get_edge_count(self) -> int:
241+
import ray
242+
208243
return ray.get(self.actor.get_edge_count.remote())
209244

210245
def get_connected_components(self, undirected: bool = True) -> List[Set[str]]:
246+
import ray
247+
211248
return ray.get(self.actor.get_connected_components.remote(undirected))
212249

213250
def has_node(self, node_id: str) -> bool:
251+
import ray
252+
214253
return ray.get(self.actor.has_node.remote(node_id))
215254

216255
def has_edge(self, source_node_id: str, target_node_id: str):
256+
import ray
257+
217258
return ray.get(self.actor.has_edge.remote(source_node_id, target_node_id))
218259

219260
def node_degree(self, node_id: str) -> int:
261+
import ray
262+
220263
return ray.get(self.actor.node_degree.remote(node_id))
221264

222265
def edge_degree(self, src_id: str, tgt_id: str) -> int:
266+
import ray
267+
223268
return ray.get(self.actor.edge_degree.remote(src_id, tgt_id))
224269

225270
def get_node(self, node_id: str) -> Any:
271+
import ray
272+
226273
return ray.get(self.actor.get_node.remote(node_id))
227274

228275
def update_node(self, node_id: str, node_data: dict[str, str]):
276+
import ray
277+
229278
return ray.get(self.actor.update_node.remote(node_id, node_data))
230279

231280
def get_all_nodes(self) -> Any:
281+
import ray
282+
232283
return ray.get(self.actor.get_all_nodes.remote())
233284

234285
def get_edge(self, source_node_id: str, target_node_id: str):
286+
import ray
287+
235288
return ray.get(self.actor.get_edge.remote(source_node_id, target_node_id))
236289

237290
def update_edge(
238291
self, source_node_id: str, target_node_id: str, edge_data: dict[str, str]
239292
):
293+
import ray
294+
240295
return ray.get(
241296
self.actor.update_edge.remote(source_node_id, target_node_id, edge_data)
242297
)
243298

244299
def get_all_edges(self) -> Any:
300+
import ray
301+
245302
return ray.get(self.actor.get_all_edges.remote())
246303

247304
def get_node_edges(self, source_node_id: str) -> Any:
305+
import ray
306+
248307
return ray.get(self.actor.get_node_edges.remote(source_node_id))
249308

250309
def upsert_node(self, node_id: str, node_data: dict[str, str]):
310+
import ray
311+
251312
return ray.get(self.actor.upsert_node.remote(node_id, node_data))
252313

253314
def upsert_edge(
254315
self, source_node_id: str, target_node_id: str, edge_data: dict[str, str]
255316
):
317+
import ray
318+
256319
return ray.get(
257320
self.actor.upsert_edge.remote(source_node_id, target_node_id, edge_data)
258321
)
259322

260323
def delete_node(self, node_id: str):
324+
import ray
325+
261326
return ray.get(self.actor.delete_node.remote(node_id))
262327

263328
def get_neighbors(self, node_id: str) -> List[str]:
329+
import ray
330+
264331
return ray.get(self.actor.get_neighbors.remote(node_id))
265332

266333
def reload(self):
334+
import ray
335+
267336
return ray.get(self.actor.reload.remote())
268337

269338

@@ -274,6 +343,8 @@ class StorageFactory:
274343

275344
@staticmethod
276345
def create_storage(backend: str, working_dir: str, namespace: str):
346+
import ray
347+
277348
if backend in ["json_kv", "rocksdb"]:
278349
actor_name = f"Actor_KV_{namespace}"
279350
actor_class = KVStorageActor

0 commit comments

Comments
 (0)