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

LLM์ฒ˜๋Ÿผ ์—ํญ๋ณด๋‹ค optimizer update ํšŸ์ˆ˜๊ฐ€ ์ค‘์š”ํ•œ ๊ฒฝ์šฐ์—๋Š” max_steps๋ฅผ ์“ด๋‹ค. max_epochs์™€ max_steps ์ค‘ ์ •ํ™•ํžˆ ํ•˜๋‚˜๋งŒ ์ง€์ •ํ•ด์•ผ ํ•œ๋‹ค.

trainer = dt.Trainer(
    max_steps=10_000,
    log_every_n_steps=50,
    val_every_n_steps=500,
    scheduler_interval="step",
    patience=4,
    monitor="val_loss",
    log_dir="runs/llm-1",
)

step ๋ชจ๋“œ์˜ ๊ธฐ๋ก์€ epoch ๋Œ€์‹  1๋ถ€ํ„ฐ ์‹œ์ž‘ํ•˜๋Š” step์„ ์“ด๋‹ค. ์œ ํ•œํ•˜๊ณ  ๋น„์–ด ์žˆ์ง€ ์•Š์€ train DataLoader๋ฅผ ๋๊นŒ์ง€ ๋Œ๋ฉด ์ž๋™์œผ๋กœ ๋‹ค์‹œ ์ˆœํšŒํ•œ๋‹ค. ๊ฒ€์ฆ์€ val_every_n_steps ๊ฐ„๊ฒฉ๊ณผ ๋งˆ์ง€๋ง‰ step์— ์‹คํ–‰ํ•˜๋ฉฐ, ์ƒ๋žตํ•˜๋ฉด ๋งˆ์ง€๋ง‰์—๋งŒ ์‹คํ–‰ํ•œ๋‹ค. best_step๊ณผ restore_best()๋„ ๊ฐ™์€ optimizer step์„ ๊ฐ€๋ฆฌํ‚จ๋‹ค.

Streamlit ์‹คํ–‰ ๋Œ€์‹œ๋ณด๋“œ

์„ ํƒ extra๋ฅผ ์„ค์น˜ํ•˜๋ฉด ๊ธฐ๋ก ์ค‘์ธ ์‹คํ–‰์„ ๋ธŒ๋ผ์šฐ์ €์—์„œ ๋น„๊ตํ•  ์ˆ˜ ์žˆ๋‹ค.

uv add "deeptool[dashboard]"
deeptool-dashboard runs/

์‹คํ–‰๋ณ„ ํ† ๊ธ€๊ณผ Focus ํŒจ๋„์„ ์ œ๊ณตํ•˜๋ฉฐ train/validation loss๋Š” ๊ฐ™์€ ์ฐจํŠธ์˜ ์‹ค์„ /์ ์„ ์œผ๋กœ ํ‘œ์‹œํ•œ๋‹ค. epoch๊ณผ step ์‹คํ–‰์€ ์„œ๋กœ ๋‹ค๋ฅธ ๊ทธ๋ฃน์ด๋‹ค. ์›๊ฒฉ SSH ํ•™์Šต์€ ์„œ๋ฒ„๋ฅผ loopback์— ๋‘” ์ฑ„ local forwarding์œผ๋กœ ๋ณธ๋‹ค.

# local
ssh -L 8501:127.0.0.1:8501 user@training-host
# remote
deeptool-dashboard runs/ --no-browser

์ž์„ธํ•œ ์‚ฌ์šฉ๋ฒ•๊ณผ 0.0.0.0 ๋…ธ์ถœ ๊ฒฝ๊ณ ๋Š” ์‹คํ–‰ ๋Œ€์‹œ๋ณด๋“œ ๊ฐ€์ด๋“œ์— ์žˆ๋‹ค.

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

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_interval="auto"๋Š” ์—ํญ ํ•™์Šต์—์„œ ์—ํญ๋งˆ๋‹ค, step ํ•™์Šต์—์„œ optimizer update๋งˆ๋‹ค ์ผ๋ฐ˜ scheduler์˜ step()์„ ๋ถ€๋ฅธ๋‹ค. ํ•„์š”ํ•˜๋ฉด "epoch" ๋˜๋Š” "step"์œผ๋กœ ๊ณ ์ •ํ•  ์ˆ˜ ์žˆ๋‹ค. ReduceLROnPlateau๋Š” ์ด ๊ฐ„๊ฒฉ๊ณผ ๋ฌด๊ด€ํ•˜๊ฒŒ ๊ฒ€์ฆ ๋’ค step(val_loss)๋กœ ํ˜ธ์ถœ๋œ๋‹ค. 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 ์„ ๋ฐ˜ํ™˜

IoUยทaccuracy์ฒ˜๋Ÿผ ํด์ˆ˜๋ก ์ข‹์€ ์ง€ํ‘œ๋Š” ๊ฒ€์ฆ ๋‹จ๊ณ„์—์„œ self.log("iou", value)๋กœ ๊ธฐ๋กํ•˜๊ณ  Trainer(monitor="iou", mode="max")๋กœ ์ง€์ •ํ•œ๋‹ค. ๊ฐ™์€ ๊ธฐ์ค€์ด best snapshot, restore_best(), ์กฐ๊ธฐ ์ข…๋ฃŒ๋ฅผ ๋ชจ๋‘ ์ œ์–ดํ•œ๋‹ค. ๊ธฐ๋ณธ์€ monitor="val_loss", mode="min"์ด๋‹ค.

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 ๋ช‡ ๋ฒˆ์˜ monitor ํ™•์ธ ๋™์•ˆ ๊ฐœ์„ ์ด ์—†์œผ๋ฉด ๋ฉˆ์ถœ ๊ฒƒ์ธ๊ฐ€
monitor "val_loss" best/์กฐ๊ธฐ ์ข…๋ฃŒ์— ์‚ฌ์šฉํ•  ์ง€ํ‘œ ์ด๋ฆ„
mode "min" min ๋˜๋Š” max

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

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_score, best_epoch, best_step, 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.4.0.tar.gz (232.2 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.4.0-py3-none-any.whl (30.1 kB view details)

Uploaded Python 3

File details

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

File metadata

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

File hashes

Hashes for deeptool-0.4.0.tar.gz
Algorithm Hash digest
SHA256 70e488bed609b9d990350795ec150a6abd61be6060e54d960cc6a2a3c1b9694e
MD5 f183a4e15ab9493a6b314b7511c29dcd
BLAKE2b-256 eb03f3c96f8f869bdad5fb9c12c797d637dbb12b75e7ade96ae29344a8028600

See more details on using hashes here.

Provenance

The following attestation bundles were made for deeptool-0.4.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.4.0-py3-none-any.whl.

File metadata

  • Download URL: deeptool-0.4.0-py3-none-any.whl
  • Upload date:
  • Size: 30.1 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.4.0-py3-none-any.whl
Algorithm Hash digest
SHA256 d8ea61285daec2029c3aba93115320c9a8211996f841fe5ef1e6ac4441f0a9e4
MD5 918a088a6d14918842ea7a6e79755151
BLAKE2b-256 dc4512ee5e7fca1ffcc5ab8787cfa63640ce00034070e1df0baa79346ee826ad

See more details on using hashes here.

Provenance

The following attestation bundles were made for deeptool-0.4.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

This release

0.4.0 This release

2 files

0.3.0

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