Views
No views yet

1import jax
2from jax import numpy as jnp
3from transformers import BertTokenizer
4from BiGS.modeling_flax_bigs import FlaxBiGSForMaskedLM
5tokenizer = BertTokenizer.from_pretrained('bert-large-uncased')
6model = FlaxBiGSForMaskedLM.from_pretrained('JunxiongWang/BiGS_4096')
7text = "The goal of life is [MASK]."
8encoded_input = tokenizer(text, return_tensors='np', padding='max_length', max_length=4096)
9output = model(**encoded_input)
10tokenizer.convert_ids_to_tokens(jnp.flip(jnp.argsort(jax.nn.softmax(output.logits[encoded_input['input_ids']==103]))[0])[:10])
11
12text = "Paris is the [MASK] of France."
13encoded_input = tokenizer(text, return_tensors='np', padding='max_length', max_length=4096)
14output = model(**encoded_input)
15tokenizer.convert_ids_to_tokens(jnp.flip(jnp.argsort(jax.nn.softmax(output.logits[encoded_input['input_ids']==103]))[0])[:10])1from BiGS.modeling_flax_bigs import FlaxBiGSForSequenceClassification
2model = FlaxBiGSForSequenceClassification.from_pretrained('JunxiongWang/BiGS_4096')1from BiGS.modeling_flax_bigs import FlaxBiGSForQuestionAnswering
2model = FlaxBiGSForQuestionAnswering.from_pretrained('JunxiongWang/BiGS_4096')1from BiGS.modeling_flax_bigs import FlaxBiGSForMultipleChoice
2model = FlaxBiGSForMultipleChoice.from_pretrained('JunxiongWang/BiGS_4096')