Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
66 changes: 0 additions & 66 deletions dooly/build_utils.py

This file was deleted.

33 changes: 13 additions & 20 deletions dooly/converters/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,10 +70,7 @@ def convert(self):
config = self.get_model_config(self._pororo_model)
hf_model = self.intialize_hf_model(config)

self._hf_model = self.porting_pororo_to_hf(
self._pororo_model,
hf_model,
)
self._hf_model = self.porting_pororo_to_hf(self._pororo_model, hf_model)
self.save_hf_model(self._hf_model)

# convert misc files
Expand Down Expand Up @@ -116,7 +113,8 @@ def load_vocab(self):
raise NotImplementedError

def save_vocab(self, vocab, filename: str = "vocab.json"):
with open(os.path.join(self.save_path, filename), "w", encoding="utf-8") as f:
save_file_path = os.path.join(self.save_path, filename)
with open(save_file_path, "w", encoding="utf-8") as f:
json.dump(vocab, f, ensure_ascii=False)

def load_and_save_vocab(self):
Expand All @@ -141,9 +139,9 @@ def get_misc_filenames(self):
def load_and_save_misc(self):
misc_files = self.get_misc_filenames()
for misc_file in misc_files:
shutil.move(
src=os.path.join(self._pororo_save_path, misc_file), dst=self.save_path
)
src_path = os.path.join(self._pororo_save_path, misc_file)
dst_path = self.save_path
shutil.move(src=src_path, dst=dst_path)


class FsmtConverter(DoolyConverter):
Expand Down Expand Up @@ -260,19 +258,14 @@ def porting_pororo_to_hf(self, pororo_model, hf_model):
sent_encoder = pororo_model.model.encoder.sentence_encoder
# Now let's copy all the weights.
# Embeddings
hf_model.roberta.embeddings.word_embeddings.weight = (
sent_encoder.embed_tokens.weight
)
hf_model.roberta.embeddings.position_embeddings.weight = (
sent_encoder.embed_positions.weight
)
hf_model.roberta.embeddings.token_type_embeddings.weight.data = (
torch.zeros_like(hf_model.roberta.embeddings.token_type_embeddings.weight)
embeddings = hf_model.roberta.embeddings
embeddings.word_embeddings.weight = sent_encoder.embed_tokens.weight
embeddings.position_embeddings.weight = sent_encoder.embed_positions.weight
embeddings.token_type_embeddings.weight.data = torch.zeros_like(
hf_model.roberta.embeddings.token_type_embeddings.weight
) # just zero them out b/c RoBERTa doesn't use them.
hf_model.roberta.embeddings.LayerNorm.weight = (
sent_encoder.emb_layer_norm.weight
)
hf_model.roberta.embeddings.LayerNorm.bias = sent_encoder.emb_layer_norm.bias
embeddings.LayerNorm.weight = sent_encoder.emb_layer_norm.weight
embeddings.LayerNorm.bias = sent_encoder.emb_layer_norm.bias

for i in range(hf_model.config.num_hidden_layers):
# Encoder: start of layer
Expand Down
6 changes: 1 addition & 5 deletions dooly/converters/kobart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,7 @@ def is_available_boto3():


class AwsS3Downloader(object):
def __init__(
self,
aws_access_key_id=None,
aws_secret_access_key=None,
):
def __init__(self, aws_access_key_id=None, aws_secret_access_key=None):
self.resource = boto3.Session(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
Expand Down
2 changes: 1 addition & 1 deletion dooly/converters/task_specific.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from .base import DoolyConverter, FsmtConverter, RobertaConverter
from .kobart_utils import download

from ..tokenizers.hf_tokenizer import (
from ..tokenizers.fast import (
convert_vocab_from_fairseq_to_hf,
build_custom_roberta_tokenizer,
PreTrainedTokenizerFast,
Expand Down
9 changes: 6 additions & 3 deletions dooly/dooly.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from typing import Optional

from .tasks import DoolyTaskHub
from .utils import _locate


def normalize_task(task: str):
Expand Down Expand Up @@ -55,7 +56,7 @@ def __new__(
raise KeyError(
f"Unavailable task name '{task}'. See here {TASK_ALIASES.keys()}"
)
task_cls = DoolyTaskHub[task]
task_cls = _locate(DoolyTaskHub[task])
if lang is not None:
lang = LANG_ALIASES.get(lang.lower(), None)
return task_cls.build(lang, n_model, **kwargs)
Expand All @@ -76,9 +77,11 @@ def available_tasks() -> str:
def available_models(task: str) -> str:
if task not in TASK_ALIASES:
raise KeyError(
f"Unknown task {task}. Please check available models via `available_tasks()`."
f"Unknown task {task}. "
"Please check available models via `available_tasks()`."
)
task_cls = DoolyTaskHub[TASK_ALIASES[normalize_task(task)]]
task_cls_path = DoolyTaskHub[TASK_ALIASES[normalize_task(task)]]
task_cls = _locate(task_cls_path)

output = f"Available models for `{task}` are "
for lang, models in task_cls.available_models.items():
Expand Down
Loading