MCPcopy Create free account
hub / github.com/baidu/DDParser / __init__

Method __init__

ddparser/run.py:281–339  ·  view source on GitHub ↗
(
        self,
        use_cuda=False,
        tree=True,
        prob=False,
        use_pos=False,
        model_files_path=None,
        buckets=False,
        batch_size=None,
        encoding_model="ernie-lstm",
    )

Source from the content-addressed store, hash-verified

279 encoding_model:指定模型,可以选lstm、transformer、ernie-1.0、ernie-tiny等
280 """
281 def __init__(
282 self,
283 use_cuda=False,
284 tree=True,
285 prob=False,
286 use_pos=False,
287 model_files_path=None,
288 buckets=False,
289 batch_size=None,
290 encoding_model="ernie-lstm",
291 ):
292 if model_files_path is None:
293 if encoding_model in ["lstm", "transformer", "ernie-1.0", "ernie-tiny", "ernie-lstm"]:
294 model_files_path = self._get_abs_path(os.path.join("./model_files/", encoding_model))
295 else:
296 raise KeyError("Unknown encoding model.")
297
298 if not os.path.exists(model_files_path):
299 try:
300 utils.download_model_from_url(model_files_path, encoding_model)
301 except Exception as e:
302 logging.error("Failed to download model, please try again")
303 logging.error("error: {}".format(e))
304 raise e
305
306 args = [
307 "--model_files={}".format(model_files_path), "--config_path={}".format(self._get_abs_path('config.ini')),
308 "--encoding_model={}".format(encoding_model)
309 ]
310
311 if use_cuda:
312 args.append("--use_cuda")
313 if tree:
314 args.append("--tree")
315 if prob:
316 args.append("--prob")
317 if batch_size:
318 args.append("--batch_size={}".format(batch_size))
319
320 args = ArgConfig(args)
321 # Don't instantiate the log handle
322 args.log_path = None
323 self.env = Environment(args)
324 self.args = self.env.args
325 paddle.set_device(self.env.place)
326 self.model = load(self.args.model_path)
327 self.model.eval()
328 self.lac = None
329 self.use_pos = use_pos
330 # buckets=None if not buckets else defaults
331 if not buckets:
332 self.args.buckets = None
333 if args.prob:
334 self.env.fields = self.env.fields._replace(PHEAD=Field("prob"))
335 if self.use_pos:
336 self.env.fields = self.env.fields._replace(CPOS=Field("postag"))
337 # set default batch size if batch_size is None and not buckets
338 if batch_size is None and not buckets:

Callers

nothing calls this directly

Calls 6

_get_abs_pathMethod · 0.95
ArgConfigClass · 0.90
EnvironmentClass · 0.90
loadFunction · 0.90
FieldClass · 0.90
evalMethod · 0.80

Tested by

no test coverage detected