(
self,
use_cuda=False,
tree=True,
prob=False,
use_pos=False,
model_files_path=None,
buckets=False,
batch_size=None,
encoding_model="ernie-lstm",
)
| 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: |
nothing calls this directly
no test coverage detected