mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Do not modify num calib data samples to batch boundary (#483)
Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
This commit is contained in:
@@ -16,7 +16,6 @@
|
||||
"""Utility functions for getting samples and forward loop function for different datasets."""
|
||||
|
||||
import copy
|
||||
import math
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from warnings import warn
|
||||
@@ -206,8 +205,6 @@ def get_dataset_dataloader(
|
||||
if isinstance(dataset_name, str):
|
||||
dataset_name = [dataset_name]
|
||||
|
||||
num_samples = [math.ceil(num_sample / batch_size) * batch_size for num_sample in num_samples]
|
||||
|
||||
assert len(dataset_name) == len(num_samples), (
|
||||
"dataset_name and num_samples must be the same length"
|
||||
)
|
||||
|
||||
@@ -15,7 +15,6 @@
|
||||
|
||||
"""Utility functions for getting samples and forward loop function for different speech datasets."""
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
@@ -101,8 +100,6 @@ def get_speech_dataset_dataloader(
|
||||
"""
|
||||
assert processor is not None, "Please provide a valid processor."
|
||||
|
||||
num_samples = math.ceil(num_samples / batch_size) * batch_size
|
||||
|
||||
dataset = _get_speech_dataset(dataset_name, num_samples=num_samples)
|
||||
first_sample = next(iter(dataset))
|
||||
first_text = first_sample["text"]
|
||||
|
||||
@@ -15,7 +15,6 @@
|
||||
|
||||
"""Utility functions for getting samples and forward loop function for different vlm datasets."""
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
@@ -93,8 +92,6 @@ def get_vlm_dataset_dataloader(
|
||||
"""
|
||||
assert processor is not None, "Please provide a valid processor."
|
||||
|
||||
num_samples = math.ceil(num_samples / batch_size) * batch_size
|
||||
|
||||
dataset = _get_vlm_dataset(dataset_name, num_samples=num_samples)
|
||||
# Apply the preprocessing function to the dataset
|
||||
processed_dataset = dataset.map(
|
||||
|
||||
Reference in New Issue
Block a user