| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,9 +1,39 @@ | |||
| 1 | 1 | import unittest | |
| 2 | 2 | ||
| 3 | + import fastai | ||
| 4 | + import pandas as pd | ||
| 5 | + import torch | ||
| 6 | + | ||
| 7 | + from fastai.docs import * | ||
| 8 | + from fastai.tabular import * | ||
| 3 | 9 | from fastai.core import partition | |
| 10 | + from fastai.torch_core import tensor | ||
| 4 | 11 | ||
| 5 | 12 | class TestFastAI(unittest.TestCase): | |
| 6 | 13 | def test_partition(self): | |
| 7 | 14 | result = partition([1,2,3,4,5], 2) | |
| 8 | 15 | ||
| 9 | 16 | self.assertEqual(3, len(result)) | |
| 17 | + | ||
| 18 | + def test_has_version(self): | ||
| 19 | + self.assertGreater(len(fastai.__version__), 1) | ||
| 20 | + | ||
| 21 | + # based on https://github.com/fastai/fastai/blob/master/tests/test_torch_core.py#L17 | ||
| 22 | + def test_torch_tensor(self): | ||
| 23 | + a = tensor([1, 2, 3]) | ||
| 24 | + b = torch.tensor([1, 2, 3]) | ||
| 25 | + | ||
| 26 | + self.assertTrue(torch.all(a == b)) | ||
| 27 | + | ||
| 28 | + def test_tabular(self): | ||
| 29 | + df = pd.read_csv("/input/tests/data/train.csv") | ||
| 30 | + | ||
| 31 | + train_df, valid_df = df[:-5].copy(),df[-5:].copy() | ||
| 32 | + dep_var = "label" | ||
| 33 | + cont_names = [] | ||
| 34 | + for i in range(784): | ||
| 35 | + cont_names.append("pixel" + str(i)) | ||
| 36 | + | ||
| 37 | + data = tabular_data_from_df("", train_df, valid_df, dep_var, cont_names=cont_names, cat_names=[]) | ||
| 38 | + learn = get_tabular_learner(data, layers=[200, 100]) | ||
| 39 | + learn.fit(epochs=1) | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,7 +6,7 @@ class TestPandas(unittest.TestCase): | |||
| 6 | 6 | def test_read_csv(self): | |
| 7 | 7 | data = pd.read_csv("/input/tests/data/train.csv") | |
| 8 | 8 | ||
| 9 | - self.assertEqual(14915, data.size) | ||
| 9 | + self.assertEqual(19, len(data.index)) | ||
| 10 | 10 | ||
| 11 | 11 | def test_read_feather(self): | |
| 12 | 12 | data = pd.read_feather("/input/tests/data/feather-0_3_1.feather") | |
| Back | FazBrowse Home | New Git URL |
0 commit comments