Skip to main content

PyTorch CRF with N-best decoding

Project description

PyTorch CRF with N-best Decoding

Implementation of Conditional Random Fields (CRF) in PyTorch 1.0. It supports top-N most probable paths decoding.

The package is based on pytorch-crf with only the following differences

  • Method _viterbi_decode that decodes the most probable path get optimized. Running time gets reduced to 50% or less with batch size 15+ and sequence length 20+
  • The class now supports decoding top-N most probable paths through the implementation of the method _viterbi_decode_nbest

Requirements

  • Python 3 (>= 3.6)
  • PyTorch 1.0

Installation

pip install pytorchcrf

Examples

>>> import torch
>>> from pytorchcrf import CRF
>>> num_tags = 4  # number of tags is 4
>>> model = CRF(num_tags)
>>> seq_length = 3  # maximum sequence length in a batch
>>> batch_size = 2  # number of samples in the batch
>>> emissions = torch.randn(seq_length, batch_size, num_tags)

>>> model.decode(emissions)
>>> model.decode(emissions, nbest=3)

Project details


Release history Release notifications

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Files for pytorchcrf, version 0.1.0
Filename, size File type Python version Upload date Hashes
Filename, size pytorchcrf-0.1.0-py3-none-any.whl (7.0 kB) File type Wheel Python version py3 Upload date Hashes View hashes

Supported by

Elastic Elastic Search Pingdom Pingdom Monitoring Google Google BigQuery Sentry Sentry Error logging AWS AWS Cloud computing DataDog DataDog Monitoring Fastly Fastly CDN SignalFx SignalFx Supporter DigiCert DigiCert EV certificate StatusPage StatusPage Status page