huggingface/Model_Cards_Writing_Tool
125
1#!/usr/bin/env python32 3import re4 5"""6Extracts code from the file "./Libraries.ts".7(Note that "Libraries.ts", must be in the same directory as 8this script).9"""10 11file = None12 13def read_file(library: str, model_name: str) -> str:14 text = file15 16 match = re.search('const ' + library + '.*', text, re.DOTALL).group()17 if match:18 text = match[match.index('`') + 1:match.index('`;')].replace('${model.id}', model_name)19 20 return text21 22file = """23import type { ModelData } from "./Types";24/**25 * Add your new library here.26 */27export enum ModelLibrary {28 "adapter-transformers" = "Adapter Transformers",29 "allennlp" = "allenNLP",30 "asteroid" = "Asteroid",31 "diffusers" = "Diffusers",32 "espnet" = "ESPnet",33 "fairseq" = "Fairseq",34 "flair" = "Flair",35 "keras" = "Keras",36 "nemo" = "NeMo",37 "pyannote-audio" = "pyannote.audio",38 "sentence-transformers" = "Sentence Transformers",39 "sklearn" = "Scikit-learn",40 "spacy" = "spaCy",41 "speechbrain" = "speechbrain",42 "tensorflowtts" = "TensorFlowTTS",43 "timm" = "Timm",44 "fastai" = "fastai",45 "transformers" = "Transformers",46 "stanza" = "Stanza",47 "fasttext" = "fastText",48 "stable-baselines3" = "Stable-Baselines3",49 "ml-agents" = "ML-Agents",50}51 52export const ALL_MODEL_LIBRARY_KEYS = Object.keys(ModelLibrary) as (keyof typeof ModelLibrary)[];53 54 55/**56 * Elements configurable by a model library.57 */58export interface LibraryUiElement {59 /**60 * Name displayed on the main61 * call-to-action button on the model page.62 */63 btnLabel: string;64 /**65 * Repo name66 */67 repoName: string;68 /**69 * URL to library's repo70 */71 repoUrl: string;72 /**73 * Code snippet displayed on model page74 */75 snippet: (model: ModelData) => string;76}77 78function nameWithoutNamespace(modelId: string): string {79 const splitted = modelId.split("/");80 return splitted.length === 1 ? splitted[0] : splitted[1];81}82 83//#region snippets84 85const adapter_transformers = (model: ModelData) =>86 `from transformers import ${model.config?.adapter_transformers?.model_class}87 88model = ${model.config?.adapter_transformers?.model_class}.from_pretrained("${model.config?.adapter_transformers?.{model.id}}")89model.load_adapter("${model.id}", source="hf")`;90 91const allennlpUnknown = (model: ModelData) =>92 `import allennlp_models93from allennlp.predictors.predictor import Predictor94 95predictor = Predictor.from_path("hf://${model.id}")`;96 97const allennlpQuestionAnswering = (model: ModelData) =>98 `import allennlp_models99from allennlp.predictors.predictor import Predictor100 101predictor = Predictor.from_path("hf://${model.id}")102predictor_input = {"passage": "My name is Wolfgang and I live in Berlin", "question": "Where do I live?"}103predictions = predictor.predict_json(predictor_input)`;104 105const allennlp = (model: ModelData) => {106 if (model.tags?.includes("question-answering")) {107 return allennlpQuestionAnswering(model);108 }109 return allennlpUnknown(model);110};111 112const asteroid = (model: ModelData) =>113 `from asteroid.models import BaseModel114 115model = BaseModel.from_pretrained("${model.id}")`;116 117const diffusers = (model: ModelData) =>118 `from diffusers import DiffusionPipeline119 120pipeline = DiffusionPipeline.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`;121 122const espnetTTS = (model: ModelData) =>123 `from espnet2.bin.tts_inference import Text2Speech124 125model = Text2Speech.from_pretrained("${model.id}")126 127speech, *_ = model("text to generate speech from")`;128 129const espnetASR = (model: ModelData) =>130 `from espnet2.bin.asr_inference import Speech2Text131 132model = Speech2Text.from_pretrained(133 "${model.id}"134)135 136speech, rate = soundfile.read("speech.wav")137text, *_ = model(speech)`;138 139const espnetUnknown = () =>140 `unknown model type (must be text-to-speech or automatic-speech-recognition)`;141 142const espnet = (model: ModelData) => {143 if (model.tags?.includes("text-to-speech")) {144 return espnetTTS(model);145 } else if (model.tags?.includes("automatic-speech-recognition")) {146 return espnetASR(model);147 }148 return espnetUnknown();149};150 151const fairseq = (model: ModelData) =>152 `from fairseq.checkpoint_utils import load_model_ensemble_and_task_from_hf_hub153 154models, cfg, task = load_model_ensemble_and_task_from_hf_hub(155 "${model.id}"156)`;157 158 159const flair = (model: ModelData) =>160 `from flair.models import SequenceTagger161 162tagger = SequenceTagger.load("${model.id}")`;163 164const keras = (model: ModelData) =>165 `from huggingface_hub import from_pretrained_keras166 167model = from_pretrained_keras("${model.id}")168`;169 170const pyannote_audio_pipeline = (model: ModelData) =>171 `from pyannote.audio import Pipeline172 173pipeline = Pipeline.from_pretrained("${model.id}")174 175# inference on the whole file176pipeline("file.wav")177 178# inference on an excerpt179from pyannote.core import Segment180excerpt = Segment(start=2.0, end=5.0)181 182from pyannote.audio import Audio183waveform, sample_rate = Audio().crop("file.wav", excerpt)184pipeline({"waveform": waveform, "sample_rate": sample_rate})`;185 186const pyannote_audio_model = (model: ModelData) =>187 `from pyannote.audio import Model, Inference188 189model = Model.from_pretrained("${model.id}")190inference = Inference(model)191 192# inference on the whole file193inference("file.wav")194 195# inference on an excerpt196from pyannote.core import Segment197excerpt = Segment(start=2.0, end=5.0)198inference.crop("file.wav", excerpt)`;199 200const pyannote_audio = (model: ModelData) => {201 if (model.tags?.includes("pyannote-audio-pipeline")) {202 return pyannote_audio_pipeline(model);203 }204 return pyannote_audio_model(model);205};206 207const tensorflowttsTextToMel = (model: ModelData) =>208 `from tensorflow_tts.inference import AutoProcessor, TFAutoModel209 210processor = AutoProcessor.from_pretrained("${model.id}")211model = TFAutoModel.from_pretrained("${model.id}")212`;213 214const tensorflowttsMelToWav = (model: ModelData) =>215 `from tensorflow_tts.inference import TFAutoModel216 217model = TFAutoModel.from_pretrained("${model.id}")218audios = model.inference(mels)219`;220 221const tensorflowttsUnknown = (model: ModelData) =>222 `from tensorflow_tts.inference import TFAutoModel223 224model = TFAutoModel.from_pretrained("${model.id}")225`;226 227const tensorflowtts = (model: ModelData) => {228 if (model.tags?.includes("text-to-mel")) {229 return tensorflowttsTextToMel(model);230 } else if (model.tags?.includes("mel-to-wav")) {231 return tensorflowttsMelToWav(model);232 }233 return tensorflowttsUnknown(model);234};235 236const timm = (model: ModelData) =>237 `import timm238 239model = timm.create_model("hf_hub:${model.id}", pretrained=True)`;240 241const sklearn = (model: ModelData) =>242 `from huggingface_hub import hf_hub_download243import joblib244 245model = joblib.load(246 hf_hub_download("${model.id}", "sklearn_model.joblib")247)`;248 249const fastai = (model: ModelData) =>250 `from huggingface_hub import from_pretrained_fastai251 252learn = from_pretrained_fastai("${model.id}")`;253 254const sentenceTransformers = (model: ModelData) =>255 `from sentence_transformers import SentenceTransformer256 257model = SentenceTransformer("${model.id}")`;258 259const spacy = (model: ModelData) =>260 `!pip install https://huggingface.co/${model.id}/resolve/main/${nameWithoutNamespace(model.id)}-any-py3-none-any.whl261 262# Using spacy.load().263import spacy264nlp = spacy.load("${nameWithoutNamespace(model.id)}")265 266# Importing as module.267import ${nameWithoutNamespace(model.id)}268nlp = ${nameWithoutNamespace(model.id)}.load()`;269 270const stanza = (model: ModelData) =>271 `import stanza272 273stanza.download("${nameWithoutNamespace(model.id).replace("stanza-", "")}")274nlp = stanza.Pipeline("${nameWithoutNamespace(model.id).replace("stanza-", "")}")`;275 276 277const speechBrainMethod = (speechbrainInterface: string) => {278 switch (speechbrainInterface) {279 case "EncoderClassifier":280 return "classify_file";281 case "EncoderDecoderASR":282 case "EncoderASR":283 return "transcribe_file";284 case "SpectralMaskEnhancement":285 return "enhance_file";286 case "SepformerSeparation":287 return "separate_file";288 default:289 return undefined;290 }291};292 293const speechbrain = (model: ModelData) => {294 const speechbrainInterface = model.config?.speechbrain?.interface;295 if (speechbrainInterface === undefined) {296 return `# interface not specified in config.json`;297 }298 299 const speechbrainMethod = speechBrainMethod(speechbrainInterface);300 if (speechbrainMethod === undefined) {301 return `# interface in config.json invalid`;302 }303 304 return `from speechbrain.pretrained import ${speechbrainInterface}305model = ${speechbrainInterface}.from_hparams(306 "${model.id}"307)308model.${speechbrainMethod}("file.wav")`;309};310 311const transformers = (model: ModelData) => {312 const info = model.transformersInfo;313 if (!info) {314 return `# ⚠️ Type of model unknown`;315 }316 if (info.processor) {317 const varName = info.processor === "AutoTokenizer" ? "tokenizer"318 : info.processor === "AutoFeatureExtractor" ? "extractor"319 : "processor"320 ;321 return [322 `from transformers import ${info.processor}, ${info.auto_model}`,323 "",324 `${varName} = ${info.processor}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,325 "",326 `model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,327 ].join("\n");328 } else {329 return [330 `from transformers import ${info.auto_model}`,331 "",332 `model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,333 ].join("\n");334 }335};336 337const fasttext = (model: ModelData) =>338 `from huggingface_hub import hf_hub_download339import fasttext340 341model = fasttext.load_model(hf_hub_download("${model.id}", "model.bin"))`;342 343const stableBaselines3 = (model: ModelData) =>344 `from huggingface_sb3 import load_from_hub345checkpoint = load_from_hub(346 repo_id="${model.id}",347 filename="{MODEL FILENAME}.zip",348)`;349 350const nemoDomainResolver = (domain: string, model: ModelData): string | undefined => {351 const modelName = `${nameWithoutNamespace(model.id)}.nemo`;352 353 switch (domain) {354 case "ASR":355 return `import nemo.collections.asr as nemo_asr356asr_model = nemo_asr.models.ASRModel.from_pretrained("${model.id}")357 358transcriptions = asr_model.transcribe(["file.wav"])`;359 default:360 return undefined;361 }362};363 364const mlAgents = (model: ModelData) =>365 `mlagents-load-from-hf --repo-id="${model.id}" --local-dir="./downloads"`;366 367const nemo = (model: ModelData) => {368 let command: string | undefined = undefined;369 // Resolve the tag to a nemo domain/sub-domain 370 if (model.tags?.includes("automatic-speech-recognition")) {371 command = nemoDomainResolver("ASR", model);372 }373 374 return command ?? `# tag did not correspond to a valid NeMo domain.`;375};376 377//#endregion378 379 380 381export const MODEL_LIBRARIES_UI_ELEMENTS: { [key in keyof typeof ModelLibrary]?: LibraryUiElement } = {382 // ^^ TODO(remove the optional ? marker when Stanza snippet is available)383 "adapter-transformers": {384 btnLabel: "Adapter Transformers",385 repoName: "adapter-transformers",386 repoUrl: "https://github.com/Adapter-Hub/adapter-transformers",387 snippet: adapter_transformers,388 },389 "allennlp": {390 btnLabel: "AllenNLP",391 repoName: "AllenNLP",392 repoUrl: "https://github.com/allenai/allennlp",393 snippet: allennlp,394 },395 "asteroid": {396 btnLabel: "Asteroid",397 repoName: "Asteroid",398 repoUrl: "https://github.com/asteroid-team/asteroid",399 snippet: asteroid,400 },401 "diffusers": {402 btnLabel: "Diffusers",403 repoName: "🤗/diffusers",404 repoUrl: "https://github.com/huggingface/diffusers",405 snippet: diffusers,406 },407 "espnet": {408 btnLabel: "ESPnet",409 repoName: "ESPnet",410 repoUrl: "https://github.com/espnet/espnet",411 snippet: espnet,412 },413 "fairseq": {414 btnLabel: "Fairseq",415 repoName: "fairseq",416 repoUrl: "https://github.com/pytorch/fairseq",417 snippet: fairseq,418 },419 "flair": {420 btnLabel: "Flair",421 repoName: "Flair",422 repoUrl: "https://github.com/flairNLP/flair",423 snippet: flair,424 },425 "keras": {426 btnLabel: "Keras",427 repoName: "Keras",428 repoUrl: "https://github.com/keras-team/keras",429 snippet: keras,430 },431 "nemo": {432 btnLabel: "NeMo",433 repoName: "NeMo",434 repoUrl: "https://github.com/NVIDIA/NeMo",435 snippet: nemo,436 },437 "pyannote-audio": {438 btnLabel: "pyannote.audio",439 repoName: "pyannote-audio",440 repoUrl: "https://github.com/pyannote/pyannote-audio",441 snippet: pyannote_audio,442 },443 "sentence-transformers": {444 btnLabel: "sentence-transformers",445 repoName: "sentence-transformers",446 repoUrl: "https://github.com/UKPLab/sentence-transformers",447 snippet: sentenceTransformers,448 },449 "sklearn": {450 btnLabel: "Scikit-learn",451 repoName: "Scikit-learn",452 repoUrl: "https://github.com/scikit-learn/scikit-learn",453 snippet: sklearn,454 },455 "fastai": {456 btnLabel: "fastai",457 repoName: "fastai",458 repoUrl: "https://github.com/fastai/fastai",459 snippet: fastai,460 },461 "spacy": {462 btnLabel: "spaCy",463 repoName: "spaCy",464 repoUrl: "https://github.com/explosion/spaCy",465 snippet: spacy,466 },467 "speechbrain": {468 btnLabel: "speechbrain",469 repoName: "speechbrain",470 repoUrl: "https://github.com/speechbrain/speechbrain",471 snippet: speechbrain,472 },473 "stanza": {474 btnLabel: "Stanza",475 repoName: "stanza",476 repoUrl: "https://github.com/stanfordnlp/stanza",477 snippet: stanza,478 },479 "tensorflowtts": {480 btnLabel: "TensorFlowTTS",481 repoName: "TensorFlowTTS",482 repoUrl: "https://github.com/TensorSpeech/TensorFlowTTS",483 snippet: tensorflowtts,484 },485 "timm": {486 btnLabel: "timm",487 repoName: "pytorch-image-models",488 repoUrl: "https://github.com/rwightman/pytorch-image-models",489 snippet: timm,490 },491 "transformers": {492 btnLabel: "Transformers",493 repoName: "🤗/transformers",494 repoUrl: "https://github.com/huggingface/transformers",495 snippet: transformers,496 },497 "fasttext": {498 btnLabel: "fastText",499 repoName: "fastText",500 repoUrl: "https://fasttext.cc/",501 snippet: fasttext,502 },503 "stable-baselines3": {504 btnLabel: "stable-baselines3",505 repoName: "stable-baselines3",506 repoUrl: "https://github.com/huggingface/huggingface_sb3",507 snippet: stableBaselines3,508 },509 "ml-agents": {510 btnLabel: "ml-agents",511 repoName: "ml-agents",512 repoUrl: "https://github.com/huggingface/ml-agents",513 snippet: mlAgents,514 },515} as const;516"""517 518 519if __name__ == '__main__':520 import sys521 library_name = "keras"522 model_name = "Distillgpt2"523 print(read_file(library_name, model_name))524 525 """"526 try:527 args = sys.argv[1:]528 if args:529 print(read_file(args[0], args[1]))530 except IndexError:531 pass532 """