diff --git a/examples/scikit-learn/configs/voting_classifier.json b/examples/scikit-learn/configs/voting_classifier.json
index 612d519..48efa90 100644
--- a/examples/scikit-learn/configs/voting_classifier.json
+++ b/examples/scikit-learn/configs/voting_classifier.json
@@ -1,8 +1,10 @@
 {
-    "@type": "sklearn-model",
-    "@model": "ensemble.VotingClassifier",
-    "estimators": [
-        ["rfc", {"@type": "sklearn-model", "@model": "ensemble.RandomForestClassifier", "n_estimators": 10}],
-        ["svc", {"@type": "sklearn-model", "@model": "svm.SVC", "gamma": "scale"}]
-    ]
+    "model": {
+        "@type": "sklearn-model",
+        "@model": "ensemble.VotingClassifier",
+        "estimators": [
+            ["rfc", {"@type": "sklearn-model", "@model": "ensemble.RandomForestClassifier", "n_estimators": 10}],
+            ["svc", {"@type": "sklearn-model", "@model": "svm.SVC", "gamma": "scale"}]
+        ]
+    }
 }
diff --git a/examples/scikit-learn/iris.py b/examples/scikit-learn/iris.py
index 1aad17f..4f74a35 100644
--- a/examples/scikit-learn/iris.py
+++ b/examples/scikit-learn/iris.py
@@ -1,9 +1,10 @@
 import typing as tp
+import copy
 import importlib
 
 import colt
 import logexp
-import sklearn
+from sklearn.base import BaseEstimator
 from sklearn.datasets import load_iris
 from sklearn.model_selection import train_test_split
 
@@ -13,29 +14,22 @@ from logger import create_logger
 logger = create_logger(__name__)
 ex = logexp.Experiment("sklearn-iris")
 
-
-@colt.register("sklearn-model")
-class SklearnModelBuilder:
-    def __init__(self, **kwargs) -> None:
-        model_path = kwargs.pop("@model")
-        self._model = self._get_model_from_sklearn(model_path)
-        self._params = kwargs
-
-    def get_model(self):
-        return self._model(**self._params)
-
-    @staticmethod
-    def _get_model_from_sklearn(model_path: str) -> tp.Any:
+@colt.register("sklearn-model", constructor="from_dict")
+class SklearnModelWrapper:
+    @classmethod
+    def from_dict(cls, model_dict: tp.Dict[str, tp.Any]) -> BaseEstimator:
+        model_path = model_dict.pop("@model")
         model_path = "sklearn." + model_path
+
         module_path, model_name = model_path.rsplit(".", 1)
 
         module = importlib.import_module(module_path)
-        model = getattr(module, model_name)
+        model_cls = getattr(module, model_name)
 
-        if isinstance(model, sklearn.base.BaseEstimator):
-            raise ValueError(f"{model_path} is not an estimator.")
+        if not issubclass(model_cls, BaseEstimator):
+            raise ValueError(f"{model_path} is not an estimator")
 
-        return model
+        return model_cls(**model_dict)
 
 
 @ex.worker("sklearn-trainer")
@@ -51,6 +45,9 @@ class TrainSklearnModel(logexp.BaseWorker):
 
     def run(self):
         logger.info("load iris dataset")
+        print(self.model)
+
+        model = colt.build(self.model)
 
         iris = load_iris()
         X, y = iris.data, iris.target
@@ -60,8 +57,6 @@ class TrainSklearnModel(logexp.BaseWorker):
 
         logger.info(f"dataset size: train={len(X_train)}, valid={len(X_valid)}")
 
-        model = colt.build(self.model).get_model()
-
         logger.info("start training")
 
         model.fit(X_train, y_train)