Skip to main content

deeptool

CI PyPI

๐Ÿ“– ๋ฌธ์„œ ยท English

์ฃผํ”ผํ„ฐ ๋…ธํŠธ๋ถ์—์„œ PyTorch ๋ชจ๋ธ์„ ๊ฐ์ฒด์ง€ํ–ฅ์œผ๋กœ ๋‹ค๋ฃจ๊ธฐ ์œ„ํ•œ ์–‡์€ ๋ณด์กฐ ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ.

๋ชจ๋ธ์€ ์œ ์ €๊ฐ€ PyTorch๋กœ ์ง์ ‘ ์ž‘์„ฑํ•œ๋‹ค. ์ด ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ๋Š” ๊ทธ ์ฃผ๋ณ€๋งŒ ๋‹ด๋‹นํ•œ๋‹ค โ€” ํ•˜์ดํผํŒŒ๋ผ๋ฏธํ„ฐ ์ž๋™ ์ €์žฅ, ์…€ ๊ฐ„ ๋ฉ”์„œ๋“œ ์ถ”๊ฐ€, ํ•™์Šต ์ค‘ ์†์‹ค ๊ณก์„  ๋ผ์ด๋ธŒ ๋ Œ๋”๋ง, ๋””๋ฐ”์ด์Šค ์ž๋™ ์„ ํƒ, ํ•™์Šต ๊ธฐ๋ก ์˜์†ํ™”, ์ฒดํฌํฌ์ธํŠธ.

์„ค์น˜

uv add deeptool     # uv ํ”„๋กœ์ ํŠธ์—
pip install deeptool
import deeptool as dt

์ด ์ €์žฅ์†Œ์—์„œ ์ง์ ‘ ๊ฐœ๋ฐœํ•˜๋ ค๋ฉด:

uv sync

ํ€ต์Šคํƒ€ํŠธ

import torch
from torch import nn
from torch.nn import functional as F

import deeptool as dt


class SyntheticRegression(dt.DataModule):
    def __init__(self, n=200, batch_size=32):
        super().__init__()
        self.save_hyperparameters()
        torch.manual_seed(0)
        self.X = torch.randn(n, 2)
        self.y = self.X @ torch.tensor([[2.0], [-3.4]]) + 4.2

    def get_dataloader(self, train):
        idx = slice(0, 160) if train else slice(160, None)
        return self.get_tensorloader((self.X, self.y), train, idx)


class LinearRegression(dt.Module):
    def __init__(self, lr=0.03):
        super().__init__()
        self.save_hyperparameters()
        self.net = nn.LazyLinear(1)

๋‹ค์Œ ์…€์—์„œ ๋ฉ”์„œ๋“œ๋ฅผ ๋ง๋ถ™์ธ๋‹ค. ํด๋ž˜์Šค๋ฅผ ๋‹ค์‹œ ์ •์˜ํ•  ํ•„์š”๊ฐ€ ์—†๋‹ค.

@dt.add_to_class(LinearRegression)
def loss(self, y_hat, y):
    return F.mse_loss(y_hat, y)


@dt.add_to_class(LinearRegression)
def configure_optimizers(self):
    return torch.optim.SGD(self.parameters(), lr=self.lr)

ํ•™์Šต์„ ๋Œ๋ฆฌ๋ฉด ์†์‹ค ๊ณก์„ ์ด ์…€ ์ถœ๋ ฅ์— ์‹ค์‹œ๊ฐ„์œผ๋กœ ๊ฐฑ์‹ ๋œ๋‹ค.

trainer = dt.Trainer(max_epochs=20)
trainer.fit(LinearRegression(), SyntheticRegression())

trainer.save_checkpoint("linreg.pt")

์ „์ฒด ์˜ˆ์ œ๋Š” examples/quickstart.ipynb ์ฐธ๊ณ .

ํ•™์Šต ๊ธฐ๋ก๊ณผ ์‹คํ–‰ ๋น„๊ต

log_dir์„ ์ฃผ๋ฉด ๋ฉ”ํƒ€๋ฐ์ดํ„ฐ์™€ ์™„๋ฃŒ๋œ ์—ํญ์„ ์ฆ‰์‹œ ๋””์Šคํฌ์— ๋‚จ๊ธด๋‹ค.

trainer = dt.Trainer(max_epochs=50, plot=False, log_dir="runs/exp1")
trainer.fit(model, data)

๋ชจ๋ธ ์•ˆ์˜ ์‚ฌ์šฉ์ž ์ง€ํ‘œ๋Š” ์ด๋ฆ„์„ ๊ทธ๋Œ€๋กœ ์“ด๋‹ค. ๊ฐ™์€ ์—ํญ์—์„œ ์—ฌ๋Ÿฌ ๋ฒˆ ๋ถ€๋ฅด๋ฉด ํ‰๊ท  ํ•œ ์ ์ด ๋œ๋‹ค.

self.log("iou", value)

์‹คํ–‰๋งˆ๋‹ค meta.json๊ณผ append-only history.jsonl์ด ์ƒ๊ธด๋‹ค. ์Šคํฌ๋ฆฝํŠธ์—์„œ๋Š” log_dir์„ ์ค€ ๊ฒฝ์šฐ์—๋งŒ ์—ํญ๋‹น ํ•œ ์ค„๋„ ์ถœ๋ ฅํ•œ๋‹ค. ๊ธฐ๋ก ์ค‘์ด๊ฑฐ๋‚˜ ์ค‘๊ฐ„์— ๋ฉˆ์ถ˜ ์‹คํ–‰๋„ ์™„๋ฃŒ๋œ ์ค„๊นŒ์ง€ ์ฝ๊ณ  ๋น„๊ตํ•  ์ˆ˜ ์žˆ๋‹ค.

runs = dt.load_runs("runs")
figures = dt.plot_runs(runs)

plot_runs๋Š” ์ง€ํ‘œ๋งˆ๋‹ค Figure ํ•˜๋‚˜๋ฅผ ๋งŒ๋“ค๊ณ  ๊ทธ ์ง€ํ‘œ๊ฐ€ ์žˆ๋Š” ์‹คํ–‰๋งŒ ๊ฒน์ณ ๊ทธ๋ฆฐ๋‹ค. ๋ชจ๋ธ๋งˆ๋‹ค ์ง€ํ‘œ ์ด๋ฆ„์ด ๋‹ฌ๋ผ๋„ ๋ณ„๋„ ์Šคํ‚ค๋งˆ๊ฐ€ ํ•„์š” ์—†๋‹ค.

ํ•™์Šต๋ฅ  ์Šค์ผ€์ค„๋Ÿฌ

์—ํญ ๊ธฐ๋ฐ˜ scheduler๋Š” optimizer์™€ ํ•จ๊ป˜ ๋ฐ˜ํ™˜ํ•œ๋‹ค.

def configure_optimizers(self):
    optim = torch.optim.Adam(self.parameters(), lr=self.lr)
    scheduler = torch.optim.lr_scheduler.StepLR(optim, step_size=10, gamma=0.1)
    return optim, scheduler

์ผ๋ฐ˜ scheduler๋Š” ์—ํญ ๋’ค step(), ReduceLROnPlateau๋Š” ๊ฒ€์ฆ ๋’ค step(val_loss)๋กœ ํ˜ธ์ถœ๋œ๋‹ค. ๋งค ๋ฐฐ์น˜ ํ˜ธ์ถœ์ด ํ•„์š”ํ•œ OneCycleLR์™€ AMP๋Š” ์•„์ง ์ง€์›ํ•˜์ง€ ์•Š๋Š”๋‹ค.

์กฐ๊ธฐ ์ข…๋ฃŒ์™€ ์ตœ์  ๊ฐ€์ค‘์น˜

๊ฐœ์„ ์ด ๋ฉˆ์ถœ ๋•Œ๊นŒ์ง€ ๋Œ๋ฆฌ๊ณ  ๊ฐ€์žฅ ์ข‹์•˜๋˜ ๊ฐ€์ค‘์น˜๋ฅผ ์“ด๋‹ค.

trainer = dt.Trainer(max_epochs=100, patience=5)
trainer.fit(model, data)

len(trainer.history["val_loss"])             # 24 โ€” 100๊นŒ์ง€ ์•ˆ ๊ฐ
trainer.best_epoch, trainer.best_val_loss    # (18, 0.2913)

trainer.restore_best()                       # 18 ์„ ๋ฐ˜ํ™˜

fit() ์€ ๊ฐ€์ค‘์น˜๋ฅผ ์ž๋™์œผ๋กœ ๋˜๋Œ๋ฆฌ์ง€ ์•Š๋Š”๋‹ค. restore_best() ๋ฅผ ๋ถ€๋ฅด๊ธฐ ์ „๊นŒ์ง€๋Š” ๋งˆ์ง€๋ง‰ epoch ์ƒํƒœ์ด๋ฏ€๋กœ ๋‘ ์‹œ์ ์˜ ์„ฑ๋Šฅ์„ ๋น„๊ตํ•  ์ˆ˜ ์žˆ๋‹ค.

๊ธฐ๋ณธ์€ ๋ฉ”๋ชจ๋ฆฌ ์Šค๋ƒ…์ƒท์ด๋‹ค. ํŒŒ์ผ๋กœ ๋‚จ๊ธฐ๋ ค๋ฉด:

dt.Trainer(max_epochs=100, patience=5, best_path="best.pt")

ํŒŒ์ผ์—๋Š” ๋ชจ๋ธ ๊ฐ€์ค‘์น˜๋งŒ ๋“ค์–ด๊ฐ„๋‹ค. optimizer ์ƒํƒœ๋Š” restore_best() ๊ฐ€ ์ฝ์ง€ ์•Š๋Š”๋ฐ Adam ๊ธฐ์ค€ ๋ชจ๋ธ์˜ 2๋ฐฐ๋ผ ๋งค epoch ์“ฐ๋ฉด ๋‚ญ๋น„๋‹ค. ์ตœ์ €์ ๋ถ€ํ„ฐ ํ•™์Šต์„ ์žฌ๊ฐœํ•  ๊ณ„ํš์ด๋ฉด best_with_optim=True ๋กœ ์ „์ฒด ์ฒดํฌํฌ์ธํŠธ๋ฅผ ๋‚จ๊ธด๋‹ค.

์ธ์ž ๊ธฐ๋ณธ ์˜๋ฏธ
snapshot_best True ์Šค๋ƒ…์ƒท์„ ๋งŒ๋“ค ๊ฒƒ์ธ๊ฐ€
best_path None None ์ด๋ฉด ๋ฉ”๋ชจ๋ฆฌ, ๊ฒฝ๋กœ๋ฉด ํŒŒ์ผ
best_with_optim False ํŒŒ์ผ์— optimizer ์ƒํƒœ๋„ ๋„ฃ์„ ๊ฒƒ์ธ๊ฐ€
patience None ๋ช‡ epoch ๊ฐœ์„ ์ด ์—†์œผ๋ฉด ๋ฉˆ์ถœ ๊ฒƒ์ธ๊ฐ€

ํ•™์Šต ํ›„ ํ‰๊ฐ€

p = trainer.predict(data)        # ๊ฒ€์ฆ์…‹ ์ „์ฒด ์ถ”๋ก 

p.accuracy                       # 0.8837
p.preds                          # ์ƒ˜ํ”Œ๋ณ„ ์˜ˆ์ธก ํด๋ž˜์Šค
p.confidence                     # ์˜ˆ์ธก ํ™•์‹ ๋„
p.correct                        # ๋งž์ท„๋Š”์ง€ ์—ฌ๋ถ€ (bool ํ…์„œ)

p = trainer.predict(data, keep_inputs=True)
p.inputs[~p.correct]             # ํ‹€๋ฆฐ ์ƒ˜ํ”Œ์˜ ์ž…๋ ฅ โ€” ์‹œ๊ฐํ™”์— ์“ด๋‹ค

predsยทprobsยทconfidenceยทcorrectยทaccuracy ๋Š” ๋ถ„๋ฅ˜ ์ „์šฉ์ด๋‹ค. ํšŒ๊ท€ ๋ชจ๋ธ์ด๋ฉด p.outputs ๋ฅผ ์ง์ ‘ ์“ด๋‹ค.

API

์ด๋ฆ„ ์—ญํ• 
dt.add_to_class(Class) ๋ฐ์ฝ”๋ ˆ์ดํŠธํ•œ ํ•จ์ˆ˜๋ฅผ Class ์˜ ๋ฉ”์„œ๋“œ๋กœ ๋“ฑ๋ก
dt.HyperParameters save_hyperparameters() ๋กœ __init__ ์ธ์ž๋ฅผ ์†์„ฑ + hparams ๋กœ ์ €์žฅ
dt.DataModule get_dataloader(train) ํ•˜๋‚˜๋งŒ ๊ตฌํ˜„ํ•˜๋ฉด ๋˜๋Š” ๋ฐ์ดํ„ฐ ๊ทœ์•ฝ
dt.Module forward/loss/configure_optimizers ๋ฅผ ์ฑ„์šฐ๋Š” ๋ชจ๋ธ ๊ทœ์•ฝ
dt.Trainer fit(model, data), predict(data), restore_best(), save_checkpoint, load_checkpoint, history, best_epoch, best_val_loss
dt.predict ๋ชจ๋ธ๊ณผ dataloader ๋ฅผ ๋ฐ›์•„ ๋ฐ์ดํ„ฐ์…‹ ์ „์ฒด ์˜ˆ์ธก์„ ๋ชจ์€๋‹ค
dt.Predictions ์˜ˆ์ธก ๊ฒฐ๊ณผ. predsยทprobsยทconfidenceยทcorrectยทaccuracy
dt.ProgressBoard ๋ผ์ด๋ธŒ ์†์‹ค ๊ณก์„ . Trainer(plot=True) ๊ฐ€ ์ž๋™์œผ๋กœ ๋งŒ๋“ ๋‹ค
dt.RunRecorder ํ•œ ์‹คํ–‰์˜ meta.json๊ณผ history.jsonl ๊ธฐ๋ก
dt.load_runs ์—ฌ๋Ÿฌ ์‹คํ–‰์˜ ์ž์œ ํ˜• JSONL ์ง€ํ‘œ๋ฅผ ๋กœ๋“œ
dt.plot_runs ์ง€ํ‘œ๋ณ„ ์‹คํ–‰ ๋น„๊ต Figure ๋ชฉ๋ก ์ƒ์„ฑ
dt.default_device() cuda โ†’ mps โ†’ cpu

๊ฐœ๋ฐœ

uv run pytest

๋ผ์ด์„ผ์Šค

MIT. LICENSE ์ฐธ๊ณ .

์„ค๊ณ„๋Š” d2l-ai/d2l-en์˜ d2l/torch.py๋ฅผ ์ฐธ๊ณ ํ–ˆ๋‹ค. ํ•ด๋‹น ์ƒ˜ํ”Œ ์ฝ”๋“œ๋Š” modified MIT(LICENSE-SAMPLECODE)๋กœ ๋ฐฐํฌ๋œ๋‹ค.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

deeptool-0.3.0.tar.gz (146.1 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

deeptool-0.3.0-py3-none-any.whl (20.9 kB view details)

Uploaded Python 3

File details

Details for the file deeptool-0.3.0.tar.gz.

File metadata

  • Download URL: deeptool-0.3.0.tar.gz
  • Upload date:
  • Size: 146.1 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for deeptool-0.3.0.tar.gz
Algorithm Hash digest
SHA256 e28882dd576c4af6886a2378ee01aea83663e4f7520374fcf05a83091d72e6fd
MD5 9d411e95e2edcd6e9a816d4297306d4d
BLAKE2b-256 68fe70e31de7f1ad57e72119f72597eb1aeb091e9cfecd2d6f58eb1491b5ad2c

See more details on using hashes here.

Provenance

The following attestation bundles were made for deeptool-0.3.0.tar.gz:

Publisher: publish.yml on sciencemj/deeptool

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file deeptool-0.3.0-py3-none-any.whl.

File metadata

  • Download URL: deeptool-0.3.0-py3-none-any.whl
  • Upload date:
  • Size: 20.9 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for deeptool-0.3.0-py3-none-any.whl
Algorithm Hash digest
SHA256 e72b9be83a6e253fd797fd4a80df30508c87ed21e2487d2dbd456217c3a78022
MD5 ba1a77b16c5a67cba3f7f7ce83de4afe
BLAKE2b-256 0d487bbc6965142582b651f6ed5d59c1a0fb2e5260af495cc59baa2e66a1fa61

See more details on using hashes here.

Provenance

The following attestation bundles were made for deeptool-0.3.0-py3-none-any.whl:

Publisher: publish.yml on sciencemj/deeptool

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

0.4.0

2 files

This release

0.3.0 This release

2 files

0.2.0

2 files

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page