FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
docker-python/tests/test_pytorch_lightning.py at master · HpcDataLab/docker-python · GitHub
HpcDataLab
/
docker-python
Public
forked from
Kaggle/docker-python
Notifications
You must be signed in to change notification settings
Fork
0
Star
0
Code
Pull requests
0
Actions
Projects
Security and quality
0
Insights
Additional navigation options
Code
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
docker-python
/
tests
/
test_pytorch_lightning.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
68 lines (48 loc) · 1.86 KB
Breadcrumbs
docker-python
/
tests
/
test_pytorch_lightning.py
Copy path
File metadata and controls
68 lines (48 loc) · 1.86 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
import
unittest
import
pytorch_lightning
as
pl
import
torch
import
torch
.
nn
.
functional
as
F
from
torch
.
utils
.
data
import
DataLoader
,
TensorDataset
class
LitDataModule
(
pl
.
LightningDataModule
):
def
__init__
(
self
,
batch_size
=
16
):
super
().
__init__
()
self
.
batch_size
=
batch_size
def
setup
(
self
,
stage
=
None
):
X_train
=
torch
.
rand
(
100
,
1
,
28
,
28
)
y_train
=
torch
.
randint
(
0
,
10
,
size
=
(
100
,))
X_valid
=
torch
.
rand
(
20
,
1
,
28
,
28
)
y_valid
=
torch
.
randint
(
0
,
10
,
size
=
(
20
,))
self
.
train_ds
=
TensorDataset
(
X_train
,
y_train
)
self
.
valid_ds
=
TensorDataset
(
X_valid
,
y_valid
)
def
train_dataloader
(
self
):
return
DataLoader
(
self
.
train_ds
,
batch_size
=
self
.
batch_size
,
shuffle
=
True
)
def
val_dataloader
(
self
):
return
DataLoader
(
self
.
valid_ds
,
batch_size
=
self
.
batch_size
,
shuffle
=
False
)
class
LitClassifier
(
pl
.
LightningModule
):
def
__init__
(
self
):
super
().
__init__
()
self
.
l1
=
torch
.
nn
.
Linear
(
28
*
28
,
10
)
def
forward
(
self
,
x
):
return
F
.
relu
(
self
.
l1
(
x
.
view
(
x
.
size
(
0
),
-
1
)))
def
training_step
(
self
,
batch
,
batch_idx
):
x
,
y
=
batch
y_hat
=
self
(
x
)
loss
=
F
.
cross_entropy
(
y_hat
,
y
)
self
.
log
(
'train_loss'
,
loss
)
return
loss
def
validation_step
(
self
,
batch
,
batch_idx
):
x
,
y
=
batch
y_hat
=
self
(
x
)
loss
=
F
.
cross_entropy
(
y_hat
,
y
)
self
.
log
(
'val_loss'
,
loss
)
def
configure_optimizers
(
self
):
return
torch
.
optim
.
Adam
(
self
.
parameters
(),
lr
=
1e-2
)
class
TestPytorchLightning
(
unittest
.
TestCase
):
def
test_version
(
self
):
self
.
assertIsNotNone
(
pl
.
__version__
)
def
test_mnist
(
self
):
dm
=
LitDataModule
()
model
=
LitClassifier
()
trainer
=
pl
.
Trainer
(
gpus
=
None
,
max_epochs
=
1
)
result
=
trainer
.
fit
(
model
,
datamodule
=
dm
)
self
.
assertTrue
(
result
)
Back
|
FazBrowse Home
|
New Git URL