Reference for ultralytics/models/yolo/pose/train.py#
This page is sourced from https://github.com/ultralytics/ultralytics/blob/main/ultralytics/models/yolo/pose/train.py. Have an improvement or example to add? Open a Pull Request — thank you! 🙏
Class ultralytics.models.yolo.pose.train.PoseTrainer#
PoseTrainer(cfg=DEFAULT_CFG, overrides: dict[str, Any] | None = None, _callbacks: dict | None = None)Bases: yolo.detect.DetectionTrainer
A class extending the DetectionTrainer class for training YOLO pose estimation models.
This trainer specializes in handling pose estimation tasks, managing model training, validation, and visualization of pose keypoints alongside bounding boxes.
Args
| Name | Type | Description | Default |
|---|---|---|---|
cfg | dict, optional | Default configuration dictionary containing training parameters. | DEFAULT_CFG |
overrides | dict, optional | Dictionary of parameter overrides for the default configuration. | None |
_callbacks | dict, optional | Dictionary of callback functions to be executed during training. | None |
Attributes
| Name | Type | Description |
|---|---|---|
args | dict | Configuration arguments for training. |
model | PoseModel | The pose estimation model being trained. |
data | dict | Dataset configuration including keypoint shape information. |
loss_names | tuple | Names of the loss components, derived from the loss dict returned by the criterion. |
Methods
| Name | Description |
|---|---|
get_dataset | Retrieve the dataset and ensure it contains the required kpt_shape key. |
get_model | Get pose estimation model with specified configuration and weights. |
get_validator | Return an instance of the PoseValidator class for validation. |
set_model_attributes | Set keypoint shape, OKS sigmas, and keypoint names on the PoseModel. |
Examples
>>> from ultralytics.models.yolo.pose import PoseTrainer
>>> args = dict(model="yolo26n-pose.pt", data="coco8-pose.yaml", epochs=3)
>>> trainer = PoseTrainer(overrides=args)
>>> trainer.train()This trainer will automatically set the task to 'pose' regardless of what is provided in overrides.
ultralytics/models/yolo/pose/train.py
class PoseTrainer(yolo.detect.DetectionTrainer):
"""A class extending the DetectionTrainer class for training YOLO pose estimation models.
This trainer specializes in handling pose estimation tasks, managing model training, validation, and visualization
of pose keypoints alongside bounding boxes.
Attributes:
args (dict): Configuration arguments for training.
model (PoseModel): The pose estimation model being trained.
data (dict): Dataset configuration including keypoint shape information.
loss_names (tuple): Names of the loss components, derived from the loss dict returned by the criterion.
Methods:
get_model: Retrieve a pose estimation model with specified configuration.
set_model_attributes: Set keypoint shape, OKS sigmas, and keypoint names on the model.
get_validator: Create a validator instance for model evaluation.
plot_training_samples: Visualize training samples with keypoints.
get_dataset: Retrieve the dataset and ensure it contains required kpt_shape key.
Examples:
>>> from ultralytics.models.yolo.pose import PoseTrainer
>>> args = dict(model="yolo26n-pose.pt", data="coco8-pose.yaml", epochs=3)
>>> trainer = PoseTrainer(overrides=args)
>>> trainer.train()
"""
def __init__(self, cfg=DEFAULT_CFG, overrides: dict[str, Any] | None = None, _callbacks: dict | None = None):
"""Initialize a PoseTrainer object for training YOLO pose estimation models.
Args:
cfg (dict, optional): Default configuration dictionary containing training parameters.
overrides (dict, optional): Dictionary of parameter overrides for the default configuration.
_callbacks (dict, optional): Dictionary of callback functions to be executed during training.
Notes:
This trainer will automatically set the task to 'pose' regardless of what is provided in overrides.
"""
if overrides is None:
overrides = {}
overrides["task"] = "pose"
super().__init__(cfg, overrides, _callbacks)Method ultralytics.models.yolo.pose.train.PoseTrainer.get_dataset#
def get_dataset(self) -> dict[str, Any]Retrieve the dataset and ensure it contains the required kpt_shape key.
Returns
| Type | Description |
|---|---|
dict | A dictionary containing the training/validation/test dataset and category names. |
Raises
| Type | Description |
|---|---|
KeyError | If the kpt_shape key is not present in the dataset. |
ultralytics/models/yolo/pose/train.py
def get_dataset(self) -> dict[str, Any]:
"""Retrieve the dataset and ensure it contains the required `kpt_shape` key.
Returns:
(dict): A dictionary containing the training/validation/test dataset and category names.
Raises:
KeyError: If the `kpt_shape` key is not present in the dataset.
"""
data = super().get_dataset()
if "kpt_shape" not in data:
raise KeyError(f"No `kpt_shape` in the {self.args.data}. See https://docs.ultralytics.com/datasets/pose")
return dataMethod ultralytics.models.yolo.pose.train.PoseTrainer.get_model#
def get_model(
self,
cfg: str | Path | dict[str, Any] | None = None,
weights: torch.nn.Module | None = None,
verbose: bool = True,
) -> PoseModelGet pose estimation model with specified configuration and weights.
Args
| Name | Type | Description | Default |
|---|---|---|---|
cfg | str | Path | dict, optional | Model configuration file path or dictionary. | None |
weights | torch.nn.Module, optional | Pretrained model whose weights are loaded into the new model. | None |
verbose | bool | Whether to display model information. | True |
Returns
| Type | Description |
|---|---|
PoseModel | Initialized pose estimation model. |
ultralytics/models/yolo/pose/train.py
def get_model(
self,
cfg: str | Path | dict[str, Any] | None = None,
weights: torch.nn.Module | None = None,
verbose: bool = True,
) -> PoseModel:
"""Get pose estimation model with specified configuration and weights.
Args:
cfg (str | Path | dict, optional): Model configuration file path or dictionary.
weights (torch.nn.Module, optional): Pretrained model whose weights are loaded into the new model.
verbose (bool): Whether to display model information.
Returns:
(PoseModel): Initialized pose estimation model.
"""
model = self.set_model_names_for_load(
PoseModel(
cfg,
nc=self.data["nc"],
ch=self.data["channels"],
data_kpt_shape=self.data["kpt_shape"],
verbose=verbose and RANK == -1,
)
)
if weights:
model.load(weights)
return modelMethod ultralytics.models.yolo.pose.train.PoseTrainer.get_validator#
def get_validator(self)Return an instance of the PoseValidator class for validation.
ultralytics/models/yolo/pose/train.py
def get_validator(self):
"""Return an instance of the PoseValidator class for validation."""
return yolo.pose.PoseValidator(
self.test_loader, save_dir=self.save_dir, args=copy(self.args), _callbacks=self.callbacks
)Method ultralytics.models.yolo.pose.train.PoseTrainer.set_model_attributes#
def set_model_attributes(self)Set keypoint shape, OKS sigmas, and keypoint names on the PoseModel.
ultralytics/models/yolo/pose/train.py
def set_model_attributes(self):
"""Set keypoint shape, OKS sigmas, and keypoint names on the PoseModel."""
super().set_model_attributes()
self.model.kpt_shape = self.data["kpt_shape"]
self.model.kpt_oks_sigmas = self.data.get("kpt_oks_sigmas")
kpt_names = self.data.get("kpt_names")
if not kpt_names:
names = list(map(str, range(self.model.kpt_shape[0])))
kpt_names = {i: names for i in range(self.model.nc)}
self.model.kpt_names = kpt_names