Knowledge Distillation in AutoMM

Open In Colab Open In SageMaker Studio Lab

Pretrained foundation models are becoming increasingly large. However, these models are difficult to deploy due to limited resources available in deployment scenarios. To benefit from large models under this constraint, you transfer the knowledge from the large-scale teacher models to the student model, with knowledge distillation. In this way, the small student model can be practically deployed under real-world scenarios, while the performance will be better than training the student model from scratch thanks to the teacher.

In this tutorial, we introduce how to adopt MultiModalPredictor for knowledge distillation. For the purpose of demonstration, we use the Question-answering NLI dataset, which comprises 104,743 question, answer pairs sampled from question answering datasets. We will demonstrate how to use a large model to guide the learning and improve the performance of a small model in AutoGluon.

Load Dataset

The Question-answering NLI dataset contains sentence pairs in English. In the label column, 0 means that the sentence is not related to the question and 1 means that the sentence is related to the question.

import datasets
from datasets import load_dataset

datasets.logging.disable_progress_bar()

dataset = load_dataset("glue", "qnli")
/home/ci/opt/venv/lib/python3.13/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html
  from .autonotebook import tqdm as notebook_tqdm
Warning: You are sending unauthenticated requests to the HF Hub. Please set a HF_TOKEN to enable higher rate limits and faster downloads.
---------------------------------------------------------------------------
HfUriError                                Traceback (most recent call last)
Cell In[2], line 6
      2 from datasets import load_dataset
      3 
      4 datasets.logging.disable_progress_bar()
      5 
----> 6 dataset = load_dataset("glue", "qnli")

File ~/opt/venv/lib/python3.13/site-packages/datasets/load.py:1667, in load_dataset(path, name, data_dir, data_files, split, cache_dir, features, download_config, download_mode, verification_mode, keep_in_memory, save_infos, revision, token, streaming, num_proc, storage_options, **config_kwargs)
   1662 verification_mode = VerificationMode(
   1663     (verification_mode or VerificationMode.BASIC_CHECKS) if not save_infos else VerificationMode.ALL_CHECKS
   1664 )
   1666 # Create a dataset builder
-> 1667 builder_instance = load_dataset_builder(
   1668     path=path,
   1669     name=name,
   1670     data_dir=data_dir,
   1671     data_files=data_files,
   1672     cache_dir=cache_dir,
   1673     features=features,
   1674     download_config=download_config,
   1675     download_mode=download_mode,
   1676     revision=revision,
   1677     token=token,
   1678     storage_options=storage_options,
   1679     **config_kwargs,
   1680 )
   1682 # Return iterable dataset in case of streaming
   1683 if streaming:

File ~/opt/venv/lib/python3.13/site-packages/datasets/load.py:1290, in load_dataset_builder(path, name, data_dir, data_files, cache_dir, features, download_config, download_mode, revision, token, storage_options, **config_kwargs)
   1288 if features is not None:
   1289     features = _fix_for_backward_compatible_features(features)
-> 1290 dataset_module = dataset_module_factory(
   1291     path,
   1292     revision=revision,
   1293     download_config=download_config,
   1294     download_mode=download_mode,
   1295     data_dir=data_dir,
   1296     data_files=data_files,
   1297     cache_dir=cache_dir,
   1298 )
   1299 # Get dataset builder class
   1300 builder_kwargs = dataset_module.builder_kwargs

File ~/opt/venv/lib/python3.13/site-packages/datasets/load.py:1182, in dataset_module_factory(path, revision, download_config, download_mode, data_dir, data_files, cache_dir, **download_kwargs)
   1177             if isinstance(e1, FileNotFoundError):
   1178                 raise FileNotFoundError(
   1179                     f"Couldn't find any data file at {relative_to_absolute_path(path)}. "
   1180                     f"Couldn't find '{path}' on the Hugging Face Hub either: {type(e1).__name__}: {e1}"
   1181                 ) from None
-> 1182             raise e1 from None
   1183 else:
   1184     raise FileNotFoundError(f"Couldn't find any data file at {relative_to_absolute_path(path)}.")

File ~/opt/venv/lib/python3.13/site-packages/datasets/load.py:1157, in dataset_module_factory(path, revision, download_config, download_mode, data_dir, data_files, cache_dir, **download_kwargs)
   1147     else:
   1148         use_exported_dataset_infos = True
   1149     return HubDatasetModuleFactory(
   1150         path,
   1151         commit_hash=commit_hash,
   1152         data_dir=data_dir,
   1153         data_files=data_files,
   1154         download_config=download_config,
   1155         download_mode=download_mode,
   1156         use_exported_dataset_infos=use_exported_dataset_infos,
-> 1157     ).get_module()
   1158 except GatedRepoError as e:
   1159     message = f"Dataset '{path}' is a gated dataset on the Hub."

File ~/opt/venv/lib/python3.13/site-packages/datasets/load.py:600, in HubDatasetModuleFactory.get_module(self)
    598     download_config.download_desc = "Downloading standalone yaml"
    599 try:
--> 600     standalone_yaml_path = cached_path(
    601         hf_dataset_url(self.name, config.REPOYAML_FILENAME, revision=self.commit_hash),
    602         download_config=download_config,
    603     )
    604     with open(standalone_yaml_path, encoding="utf-8") as f:
    605         standalone_yaml_data = yaml.safe_load(f.read())

File ~/opt/venv/lib/python3.13/site-packages/datasets/utils/file_utils.py:175, in cached_path(url_or_filename, download_config, **download_kwargs)
    169 # Download files from Hugging Face.
    170 # Note: no need to check for https://huggingface.co file URLs since _prepare_path_and_storage_options
    171 # prepares Hugging Face HTTP URLs as hf:// paths already
    172 if url_or_filename.startswith("hf://") and not url_or_filename.startswith("hf://buckets/"):
    173     resolved_path = huggingface_hub.HfFileSystem(
    174         endpoint=config.HF_ENDPOINT, token=download_config.token
--> 175     ).resolve_path(url_or_filename)
    176     try:
    177         output_path = huggingface_hub.HfApi(
    178             endpoint=config.HF_ENDPOINT,
    179             token=download_config.token,
   (...)    188             force_download=download_config.force_download,
    189         )

File ~/opt/venv/lib/python3.13/site-packages/huggingface_hub/hf_file_system.py:307, in HfFileSystem.resolve_path(self, path, revision)
    302 if path.count("/") == 0:
    303     raise ValueError(
    304         f"Repository id must be 'namespace/name', got '{path}'. Single-segment ids (e.g. 'gpt2') are no longer supported."
    305     )
--> 307 parsed = parse_hf_uri(f"{constants.HF_PROTOCOL}{path}")
    309 # --- Buckets ---
    310 if parsed.is_bucket:

File ~/opt/venv/lib/python3.13/site-packages/huggingface_hub/utils/_hf_uris.py:319, in parse_hf_uri(uri, endpoint)
    317 if type_ == "bucket":
    318     return _parse_bucket_body(location, type_, raw=raw)
--> 319 return _parse_repo_body(location, type_, raw=raw)

File ~/opt/venv/lib/python3.13/site-packages/huggingface_hub/utils/_hf_uris.py:617, in _parse_repo_body(location, type_, raw)
    615     raise HfUriError(uri=raw, msg="Missing repository id before '@'.")
    616 if repo_id.count("/") != 1:
--> 617     raise HfUriError(uri=raw, msg=f"Repository id must be 'namespace/name', got '{repo_id}'.")
    618 # Special refs like 'refs/pr/10' contain '/' and must be matched eagerly,
    619 # otherwise we would split them at the first '/' and treat the rest as a path.
    620 match = _SPECIAL_REFS_REVISION_REGEX.search(rev_and_path)

HfUriError: Invalid HF URI 'hf://datasets/glue@bcdcba79d07bc864c1c254ccfcedcce55bcc9a8c/.huggingface.yaml'. Repository id must be 'namespace/name', got 'glue'.
dataset['train']
from sklearn.model_selection import train_test_split

train_valid_df = dataset["train"].to_pandas()[["question", "sentence", "label"]].sample(1000, random_state=123)
train_df, valid_df = train_test_split(train_valid_df, test_size=0.2, random_state=123)
test_df = dataset["validation"].to_pandas()[["question", "sentence", "label"]].sample(1000, random_state=123)

Load the Teacher Model

In our example, we will directly load a teacher model with the google/bert_uncased_L-12_H-768_A-12 backbone that has been trained on QNLI and distill it into a student model with the google/bert_uncased_L-6_H-768_A-12 backbone.

!wget --quiet https://automl-mm-bench.s3.amazonaws.com/unit-tests/distillation_sample_teacher.zip -O distillation_sample_teacher.zip
!unzip -q -o distillation_sample_teacher.zip -d .
from autogluon.multimodal import MultiModalPredictor

teacher_predictor = MultiModalPredictor.load("ag_distillation_sample_teacher/")

Distill to Student

Training the student model is straight forward. You may just add the teacher_predictor argument when calling .fit(). Internally, the student will be trained by matching the prediction/feature map from the teacher. It can perform better than directly finetuning the student.

student_predictor = MultiModalPredictor(label="label")
student_predictor.fit(
    train_df,
    tuning_data=valid_df,
    teacher_predictor=teacher_predictor,
    hyperparameters={
        "model.hf_text.checkpoint_name": "google/bert_uncased_L-6_H-768_A-12",
        "optim.max_epochs": 2,
    }
)
print(student_predictor.evaluate(data=test_df))

More about Knowledge Distillation

To learn how to customize distillation and how it compares with direct finetuning, see the distillation examples and README in AutoMM Distillation Examples. Especially the multilingual distillation example with more details and customization.

Other Examples

You may go to AutoMM Examples to explore other examples about AutoMM.

Customization

To learn how to customize AutoMM, please refer to Customize AutoMM.