modelmark 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,162 @@
1
+ from __future__ import annotations
2
+
3
+ import numpy as np
4
+ import time
5
+ import torch
6
+ import torch.nn as nn
7
+ from thop import profile
8
+ from torch.utils.data import DataLoader
9
+ from torch.nn.utils import clip_grad_norm_
10
+ from modelmark.common.loader import Loader
11
+
12
+ import logging
13
+ logger = logging.getLogger(__name__)
14
+ from modelmark.common.utils import load_config
15
+ config = load_config()
16
+
17
+ class Tester:
18
+ """This class runs a benchmark of the current model and loader"""
19
+
20
+ def __init__(self, model : nn.Module, loader : Loader):
21
+
22
+ self.model = model
23
+ self.train_loader = loader.train_loader
24
+ self.val_loader = loader.val_loader
25
+ self.test_loader = loader.test_loader
26
+
27
+ self.optimizer = config.test_config["optim"]["class"](self.model.parameters(), lr=config.test_config["optim"]["lr"])
28
+ self.criterion = config.test_config["optim"]["criterion"]
29
+ self.metrics = [v for _, v in config.test_config["metrics"].items()]
30
+ self.num_epochs = config.test_config["optim"]["num_epochs"]
31
+ self.max_norm = config.test_config["optim"]["max_norm"]
32
+
33
+ self.device = config.test_config["optim"]["device"]
34
+
35
+ def test(self):
36
+ """
37
+ Train, validate, return to the best model weight state according to the validation loss, test and return.
38
+ """
39
+
40
+ best_val_loss = float("inf")
41
+ best_state = None
42
+ train_time, train_gflops, train_mem = None, None, None
43
+
44
+ for epoch in range(self.num_epochs):
45
+
46
+ results = self.train(self.train_loader, measure_time=True, measure_flops=True, measure_memory=True)
47
+
48
+ train_loss = results[0]
49
+ train_time = results[1]
50
+ train_gflops = results[2]
51
+ train_mem = results[3]
52
+
53
+ val_loss = self.evaluate(self.val_loader, self.criterion)
54
+
55
+ logger.debug(f"Epoch {epoch:03d} | T {train_loss:.4f} | V {val_loss:.4f}")
56
+
57
+ # Save best model according to validation performance
58
+ if val_loss < best_val_loss:
59
+ best_val_loss = val_loss
60
+
61
+ best_state = {
62
+ k: v.detach().cpu().clone()
63
+ for k, v in self.model.state_dict().items()
64
+ }
65
+
66
+ # Restore best checkpoint
67
+ self.model.load_state_dict(best_state)
68
+
69
+ # Get final test metrics values
70
+ test_metrics = np.array([self.evaluate(self.test_loader, m) for m in self.metrics])
71
+
72
+ return test_metrics, train_time, train_gflops, train_mem
73
+
74
+ def train(self, loader: DataLoader, measure_time: bool = False,
75
+ measure_flops: bool = False, measure_memory: bool = False):
76
+ """Train the model for one epoch."""
77
+
78
+ self.model.train()
79
+
80
+ total_loss = 0.0
81
+ total_samples = 0
82
+ total_time = 0
83
+ total_flops = 0
84
+
85
+ is_cuda = self.device.type == "cuda"
86
+
87
+ if measure_memory and is_cuda:
88
+ torch.cuda.reset_peak_memory_stats(self.device)
89
+
90
+ if measure_time:
91
+ if is_cuda:
92
+ torch.cuda.synchronize()
93
+ start_time = time.perf_counter()
94
+
95
+ for i, (x, y) in enumerate(loader):
96
+
97
+ x = x.to(self.device)
98
+ y = y.to(self.device)
99
+
100
+ if measure_flops and (i == 0):
101
+ with torch.no_grad():
102
+ flops_per_batch, _ = profile(self.model, inputs=(x,), verbose=False)
103
+ flops_per_batch *= 3
104
+
105
+ self.optimizer.zero_grad(set_to_none=True)
106
+ pred = self.model(x)
107
+ loss = self.criterion(pred.squeeze(), y.squeeze())
108
+ loss.backward()
109
+ clip_grad_norm_(self.model.parameters(), max_norm=self.max_norm)
110
+ self.optimizer.step()
111
+
112
+ batch_size = x.size(0)
113
+ total_loss += loss.item() * batch_size
114
+ total_samples += batch_size
115
+
116
+ if measure_flops:
117
+ total_flops = flops_per_batch
118
+
119
+ avg_loss = total_loss / total_samples
120
+
121
+ results = [avg_loss]
122
+
123
+ if measure_time:
124
+ if is_cuda:
125
+ torch.cuda.synchronize()
126
+ total_time = time.perf_counter() - start_time
127
+ results.append(total_time)
128
+
129
+ if measure_flops:
130
+ results.append(total_flops / 1e9)
131
+
132
+ if measure_memory:
133
+ if is_cuda:
134
+ peak_mem_gb = torch.cuda.max_memory_allocated(self.device) / (1024 ** 3)
135
+ else:
136
+ peak_mem_gb = 0.0 # or float('nan') to signal "not applicable"
137
+ results.append(peak_mem_gb)
138
+
139
+ return tuple(results)
140
+
141
+ @torch.no_grad()
142
+ def evaluate(self, loader: DataLoader, criterion):
143
+ """Evaluate the model on provided data"""
144
+
145
+ self.model.eval()
146
+
147
+ total_loss = 0.0
148
+ total_samples = 0
149
+
150
+ for x, y in loader:
151
+ x = x.to(self.device)
152
+ y = y.to(self.device)
153
+
154
+ pred = self.model(x)
155
+ loss = criterion(pred.squeeze(), y.squeeze())
156
+
157
+ batch_size = x.size(0)
158
+ total_loss += loss.item() * batch_size
159
+ total_samples += batch_size
160
+
161
+ return total_loss / total_samples
162
+
@@ -0,0 +1,105 @@
1
+ from __future__ import annotations
2
+
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+
7
+ import os
8
+ import sys
9
+ import time
10
+ import random
11
+ import numpy as np
12
+
13
+ import importlib.util
14
+ import modelmark.constants as constants
15
+
16
+ import logging
17
+ logger = logging.getLogger(__name__)
18
+
19
+ def set_seed(s : int = 0):
20
+ # Seed all the generators for reproducible results
21
+ torch.random.manual_seed(s)
22
+ torch.cuda.manual_seed(s)
23
+ torch.cuda.manual_seed_all(s)
24
+ torch.use_deterministic_algorithms(True)
25
+ torch.backends.cudnn.deterministic = True
26
+ torch.backends.cudnn.benchmark = False
27
+ np.random.seed(s)
28
+ random.seed(s)
29
+
30
+ def count_parameters(module: nn.Module) -> int:
31
+ """Count the total number of elements/parameters in any PyTorch module."""
32
+ return sum(p.numel() for p in module.parameters()) / 1000
33
+
34
+ # Setup non-blocking key input based on Operating System
35
+ if os.name == 'nt':
36
+ import msvcrt
37
+ def get_keypress():
38
+ if msvcrt.kbhit():
39
+ # Return decoded string character
40
+ return msvcrt.getch().decode('utf-8', errors='ignore')
41
+ return None
42
+ else:
43
+ import select
44
+ import termios
45
+ import tty
46
+ def get_keypress():
47
+ # Check if stdin has data waiting
48
+ if select.select([sys.stdin], [], [], 0)[0]:
49
+ fd = sys.stdin.fileno()
50
+ old_settings = termios.tcgetattr(fd)
51
+ try:
52
+ tty.setraw(sys.stdin.fileno())
53
+ ch = sys.stdin.read(1)
54
+ finally:
55
+ termios.tcsetattr(fd, termios.TCSADRAIN, old_settings)
56
+ return ch
57
+ return None
58
+
59
+ def start_timer(duration_seconds):
60
+ start_time = time.time()
61
+ #print("Timer started. Press any key to interrupt...")
62
+
63
+ while True:
64
+ elapsed = time.time() - start_time
65
+ remaining = max(0, duration_seconds - elapsed)
66
+
67
+ sys.stdout.write(f"\rStart in: {remaining:.1f}s")
68
+ sys.stdout.flush()
69
+
70
+ # Check for user interruption
71
+ pressed_key = get_keypress()
72
+ if pressed_key is not None:
73
+ print("\n")
74
+ return pressed_key
75
+
76
+ if remaining <= 0:
77
+ break
78
+
79
+ time.sleep(0.05) # Lower sleep window for snappier key detection
80
+
81
+ print("\n")
82
+ #print("\n\nTimer finished naturally!")
83
+ return None
84
+
85
+ def load_config():
86
+ """Load the config (user or pkg)"""
87
+
88
+ user_config_path = constants.USER_CONFIG_PATH
89
+
90
+ if user_config_path.exists():
91
+
92
+ # Add the user working dir to sys.path
93
+ user_dir = str(constants.USER_DIR)
94
+ if user_dir not in sys.path:
95
+ sys.path.insert(0, user_dir)
96
+
97
+ spec = importlib.util.spec_from_file_location("config", user_config_path)
98
+ user_config = importlib.util.module_from_spec(spec)
99
+ spec.loader.exec_module(user_config)
100
+ return user_config
101
+
102
+ else:
103
+
104
+ from modelmark import config
105
+ return config
modelmark/config.py ADDED
@@ -0,0 +1,97 @@
1
+ """This is the testing config file example, customize as you need."""
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+
6
+ # ---- Model import secion ---- #
7
+ # Put your model class definitions here
8
+ # Either built-in:
9
+ from modelmark.models.gru import GRUModel
10
+ from modelmark.models.conv import ConvModel
11
+ from modelmark.models.lstm import LSTMModel
12
+ from modelmark.models.linear import Linear
13
+ # Either custom (uncomment)
14
+ #from models.linear import Linear
15
+
16
+ # ---- Model configuration section ----
17
+ model_config = {
18
+ "Linear" : {
19
+ "num_layers": 3,
20
+ "hidden_size": 64,
21
+ "class": Linear
22
+ },
23
+ "Conv" : {
24
+ "num_layers": 3,
25
+ "hidden_size": 64,
26
+ "kernel_size": 4,
27
+ "class": ConvModel,
28
+ },
29
+ "GRU" : {
30
+ "num_layers": 2,
31
+ "hidden_size": 64,
32
+ "class": GRUModel,
33
+ },
34
+ "LSTM" : {
35
+ "num_layers": 2,
36
+ "hidden_size": 64,
37
+ "class": LSTMModel,
38
+ },
39
+ }
40
+
41
+ # ---- Data configuration section ----
42
+ # Your datasets folder path
43
+ data_path = "data/"
44
+ # Configuration of your datasets
45
+ data_config = {
46
+ "ETTh1" : {
47
+ "path": "ett/ETTh1.csv",
48
+ "input_features": ["HUFL", "HULL", "MUFL", "MULL", "LUFL", "LULL", "OT"],
49
+ "output_features": ["OT"],
50
+ "train_ratio": 0.6,
51
+ "val_ratio": 0.2,
52
+ },
53
+ "ETTh2" : {
54
+ "path": "ett/ETTh2.csv",
55
+ "input_features": ["HUFL", "HULL", "MUFL", "MULL", "LUFL", "LULL", "OT"],
56
+ "output_features": ["OT"],
57
+ "train_ratio": 0.6,
58
+ "val_ratio": 0.2,
59
+ },
60
+ }
61
+
62
+ # ---- Testing configuration section ----
63
+
64
+ # You can define your own metric
65
+ import torch
66
+ import torch.nn.functional as F
67
+ def rmse(x, y):
68
+ return torch.sqrt(F.mse_loss(x, y))
69
+ from torch.optim import Adam
70
+
71
+ test_config = {
72
+ "optim": {
73
+ "name" : "Adam",
74
+ "class": Adam,
75
+ "max_norm": 1.0,
76
+ "criterion": nn.MSELoss(),
77
+ "lr": 1e-3,
78
+ "batch_size": 512,
79
+ "num_epochs": 5,
80
+ "device": torch.device("cuda"),
81
+ },
82
+ "metrics":
83
+ {
84
+ "RMSE": rmse,
85
+ "MAE": nn.L1Loss(),
86
+ },
87
+ "seeds": [1, 2, 3],
88
+ "contexts": [96, 192, 336, 720],
89
+ }
90
+
91
+ # Other
92
+ _data_split = ", ".join([f"{k} ({v["train_ratio"]:.2f}/{v["val_ratio"]:.2f}/{(1 - (v["train_ratio"] + v["val_ratio"])):.2f})" for k, v in data_config.items()])
93
+ # --- Report configuration section ---
94
+ title = f"Models evaluation report."
95
+ subtitle=f"""Models were trained with {test_config["optim"]["name"]} optimizer, batch size = {test_config["optim"]["batch_size"]}, LR = {test_config["optim"]["lr"]}, epochs = {test_config["optim"]["num_epochs"]}
96
+ Dataset split (train/val/test): {_data_split}
97
+ Results are mean values over {len(test_config["seeds"])} runs, seeds used: {", ".join(map(str, test_config["seeds"]))}"""
modelmark/constants.py ADDED
@@ -0,0 +1,34 @@
1
+ from pathlib import Path
2
+
3
+ # ModelMark pkg dir path
4
+ PACKAGE_DIR = Path(__file__).parent
5
+ # User cwd path
6
+ USER_DIR = Path.cwd()
7
+
8
+ # Create user config dir (if not already exists)
9
+ USER_CONFIG_DIR = USER_DIR / "modelmark_files"
10
+
11
+ # Pkg config file path
12
+ PACKAGE_CONFIG_PATH = PACKAGE_DIR / "config.py"
13
+ # User config file path
14
+ USER_CONFIG_PATH = USER_CONFIG_DIR / "config.py"
15
+
16
+ # Create user models dir (if not already created)
17
+ USER_MODELS_DIR = USER_DIR / "models"
18
+
19
+ # Pkg example model path
20
+ PACKAGE_MODEL_PATH = PACKAGE_DIR / "models" / "linear.py"
21
+ # User example mopdel path
22
+ USER_MODEL_PATH = USER_DIR / "models" / "linear.py"
23
+
24
+ PARSER_DESC = """MODELMARK
25
+
26
+ Python tool to measure the performance of a custom neural network model
27
+ and compare it to other popular architectures."""
28
+ PARSER_HELP = """\nModelmark available TASKs:
29
+ init - Test init, create the modelmark_files/, customizable config.py file and models/ with examples.
30
+ load - Dataset download, create the data/ folder and download ETT dataset there.
31
+ run - Run the test.
32
+ clear - Remove the modelmark_files folder.
33
+ reset - Runs clean and init.
34
+ """
modelmark/modelmark.py ADDED
@@ -0,0 +1,206 @@
1
+ """
2
+ ModelMark
3
+
4
+ This tool will help you to test NN models against each other,
5
+ and form a detailed report that is easy to embed to a website.
6
+
7
+ It can download dataset with
8
+ modelmark -t load
9
+ It can run the test with
10
+ modelmark -t run
11
+
12
+ Test consists of F * C * M * S steps, where:
13
+ F - number of dataset files in the config (e.g. ["ETTh1" : ..., "Weather" : ...] means F = 2)
14
+ C - number of context sizes (e.g. [32, 64, 128] means C = 3)
15
+ M - number of models (e.g. ["Linear" : ..., "LSTM" : ...] means M = 2)
16
+ S - number of seeds (e.g. [42, 43, 44] means S = 2)
17
+
18
+ At each testing run iteration, modelmark:
19
+
20
+ 1) Selects next dataset, context, model and seed
21
+ 2) Seeds the generators for reproducibility
22
+ 3) Creates the loader, model and tester objects
23
+ 4) Trains the model for E epochs, restores the state with the least validation loss
24
+ 5) Tracks the GFLOPs (AVG over one batch), Memory (AVG peak usage during full training), Time (AVG per epoch)
25
+ 5) Evaluates the model on metrics from configuration file (config.test_metric1&2)
26
+ 6) Stores the mean result over S runs
27
+
28
+ That way, the more seeds you run, the more "fair" the results are.
29
+
30
+ Finally, modelmark will form the report with all the testing results, training stats and your machine metadata.
31
+
32
+ About training stats:
33
+
34
+ Time - average time per epoch
35
+ Params - total number of model params
36
+ GFLOPs - average per batch
37
+ Peak Memory - max per training iteration
38
+
39
+ ---
40
+ Example of the model: "modelmark/models/linear.py"
41
+ Example of the config: "modelmark/config.py"
42
+ """
43
+
44
+ import sys
45
+ import torch
46
+ import numpy as np
47
+ import pandas as pd
48
+ from tqdm import tqdm
49
+
50
+ # Import custom functions
51
+ from modelmark.common.utils import set_seed, count_parameters, start_timer
52
+ from modelmark.common.logger import setup_logging
53
+ # Import custom classes
54
+ from modelmark.common.report import Report
55
+ from modelmark.common.parser import Parser
56
+ from modelmark.common.loader import Loader
57
+ from modelmark.common.tester import Tester
58
+
59
+ import logging
60
+ logger = logging.getLogger(__name__)
61
+ from rich.console import Console
62
+ console = Console()
63
+ from modelmark.common.utils import load_config
64
+ config = load_config()
65
+
66
+ def _check_model(name = "Linear", device=config.test_config["optim"]["device"]):
67
+ """Check the testing model on shapes compatibility."""
68
+ first_config = next(iter(config.data_config.values()))
69
+
70
+ input_dim = len(first_config["input_features"])
71
+ output_dim = len(first_config["output_features"])
72
+
73
+ model = config.model_config[name]["class"](input_dim, output_dim, context_size = 10).to(device)
74
+ dummy_input = torch.zeros([4, 10, input_dim]).to(device)
75
+ dummy_target = torch.zeros([4, 10, output_dim]).to(device)
76
+ generated = model(dummy_input)
77
+ if generated.shape == dummy_target.shape:
78
+ console.print(f"Testing model {name} is ready", style="green")
79
+ return None
80
+ else:
81
+ console.print(f"Testing model output expected shape {dummy_target.shape}, got {generated.shape}.\nCheck your definition of {name} model.", style="red")
82
+ return 1
83
+
84
+ def run() -> int:
85
+
86
+ # Parse the arguments
87
+ parser = Parser()
88
+ code = parser.run()
89
+ if code is not None:
90
+ return code
91
+
92
+ # Setup logger
93
+ setup_logging()
94
+ logger.info("Logger started")
95
+
96
+ # Check the models
97
+ for k, _ in config.model_config.items():
98
+ if _check_model(k) is not None:
99
+ return 1
100
+
101
+ # Display test information
102
+ logger.info("Display test info.")
103
+ console.print("-" * 128, style = "cyan")
104
+ console.print("Info\n", style = "magenta")
105
+ console.print(f"This test runs multiple training/validation/testing passes of all configured model architectures: {[k for k, _ in config.model_config.items()]}", style = "white")
106
+ console.print(f"To provide you with the most fair comparision, we run multiple tests over seeds, and average the results.", style = "white")
107
+ console.print(f"There is no parallel execution at the moment, so the test may take a while to finish (it depends on your config and hardware).", style = "white")
108
+ console.print("-" * 128, style="cyan")
109
+
110
+ console.print("Start the test? [y/n] (Will start automatically in 30s.)", style = "white")
111
+ k = start_timer(30)
112
+ if (k == 'n') or (k == 'N'):
113
+ logger.error("Test terminated by user.")
114
+ console.print("Terminated", style="red")
115
+ return 1
116
+
117
+ # Start the test
118
+ logger.info("Starting the test.")
119
+ console.print("Starting the test.", style = "green")
120
+
121
+ # Create test buffers
122
+ report_data = []
123
+ report_dataset_names = list(config.data_config)
124
+ report_model_names = list(config.model_config)
125
+ report_metric_names = list(config.test_config["metrics"])
126
+ report_stats = np.zeros((len(report_model_names), 3))
127
+ report_params = []
128
+
129
+ device = config.test_config["optim"]["device"]
130
+ c_current, c_total = 1, len(report_dataset_names) * len(report_model_names) * len(config.test_config["seeds"]) * len(config.test_config["contexts"])
131
+ pbar = tqdm(total=c_total, desc="Progress")
132
+ for file_name, file_config in config.data_config.items():
133
+
134
+ for test_context in config.test_config["contexts"]:
135
+
136
+ report_line = []
137
+ for model_name, model_config in config.model_config.items():
138
+
139
+ total_test_metrics = np.zeros(len(config.test_config["metrics"]))
140
+ for test_seed in config.test_config["seeds"]:
141
+
142
+ # (Re-)Seed
143
+ set_seed(test_seed)
144
+ tqdm.write(f"Running test ({c_current}/{c_total}) | File: {file_name} | Context: {test_context} | Model: {model_name} | Seed: {test_seed}")
145
+ logger.debug(f"Test ({c_current}/{c_total})| F={file_name} C={test_context} M={model_name} S={test_seed}")
146
+
147
+ # Create loader
148
+ loader = Loader(file_config = file_config, context_size = test_context)
149
+ # Create model
150
+ input_dim = len(file_config["input_features"])
151
+ output_dim = len(file_config["output_features"])
152
+ model = model_config["class"](input_dim, output_dim, context_size = test_context).to(device)
153
+ # Create tester
154
+ tester = Tester(model = model, loader = loader)
155
+
156
+ # Run the test
157
+ test_metrics, train_time, train_gflops, train_mem = tester.test()
158
+ logger.debug(f"test_metrics={test_metrics} time={train_time} gflops={train_gflops} mem={train_mem}")
159
+
160
+ # Accomulate
161
+ total_test_metrics += test_metrics
162
+ model_id = report_model_names.index(model_name)
163
+ report_stats[model_id] += np.array([train_time, train_gflops, train_mem])
164
+
165
+ # Update the progress bar
166
+ c_current += 1
167
+ pbar.update(1)
168
+
169
+ # Calculate the stats (average over seeds)
170
+ test_metrics_mean = total_test_metrics / len(config.test_config["seeds"])
171
+ if len(report_params) < len(report_model_names):
172
+ report_params.append(count_parameters(model))
173
+ # Add to the current "line"
174
+ for m in test_metrics_mean:
175
+ report_line.append(f"{m:.3f}")
176
+ # Add line to the final report
177
+ report_data.append(report_line)
178
+ pbar.close()
179
+
180
+ # Calculate the report stats
181
+ report_stats = report_stats.T
182
+ report_stats = report_stats / c_total
183
+
184
+ report_time = [f"{v:.3f}" for v in report_stats[0]]
185
+ report_gflops = [f"{v:.2f}" for v in report_stats[1]]
186
+ report_mem = [f"{v:.2f}" for v in report_stats[2]]
187
+ report_params = [f"{v:.2f}" for v in report_params ]
188
+ # Obtain the columns&rows names
189
+
190
+ report_columns = pd.MultiIndex.from_product([report_model_names, report_metric_names], names=["Model", "Metric"])
191
+ report_rows = pd.MultiIndex.from_product([report_dataset_names, config.test_config["contexts"]])
192
+ # Pack to the dataframe
193
+ df = pd.DataFrame(data = report_data, index = report_rows, columns = report_columns)
194
+
195
+ model_stats={
196
+ "Time (s/epoch)": report_time,
197
+ "Params count (K)": report_params,
198
+ "GFLOPs (f/batch)": report_gflops,
199
+ "Peak Memory (Gb)": report_mem
200
+ }
201
+
202
+ report = Report()
203
+ # Form and save the report
204
+ report.report(df, config.title, config.subtitle, model_stats)
205
+
206
+ return 0
@@ -0,0 +1,75 @@
1
+ import torch
2
+ import torch.nn as nn
3
+
4
+ from modelmark.common.utils import load_config
5
+
6
+ class ConvModel(nn.Module):
7
+ def __init__(self, input_dim : int, output_dim : int, context_size : int):
8
+ super().__init__()
9
+
10
+ # To avoid circular import, now it lives here
11
+ config = load_config()
12
+
13
+ self.input_size = input_dim
14
+ self.output_size = output_dim
15
+ self.context_size = context_size
16
+
17
+ self.hidden_size = config.model_config["Conv"]["hidden_size"]
18
+ self.num_layers = config.model_config["Conv"]["num_layers"]
19
+ self.kernel_size = config.model_config["Conv"]["kernel_size"]
20
+
21
+ layers = []
22
+
23
+ in_channels = self.input_size
24
+
25
+ for _ in range(self.num_layers):
26
+ layers.append(
27
+ nn.Conv1d(
28
+ in_channels=in_channels,
29
+ out_channels=self.hidden_size,
30
+ kernel_size=self.kernel_size,
31
+ padding=self.kernel_size // 2,
32
+ )
33
+ )
34
+ layers.append(nn.ReLU())
35
+
36
+ in_channels = self.hidden_size
37
+
38
+ self.conv = nn.Sequential(*layers)
39
+
40
+ # Collapse the sequence dimension to one feature vector
41
+ self.pool = nn.AdaptiveAvgPool1d(1)
42
+
43
+ # Produce exactly context_size * output_size values
44
+ self.fc = nn.Linear(
45
+ self.hidden_size,
46
+ context_size * self.output_size,
47
+ )
48
+
49
+ def forward(self, x):
50
+ # x: [batch, sequence_length, input_size]
51
+
52
+ # Conv1d expects:
53
+ # [batch, channels, sequence_length]
54
+ x = x.transpose(1, 2)
55
+
56
+ # [batch, hidden_size, sequence_length]
57
+ x = self.conv(x)
58
+
59
+ # [batch, hidden_size, 1]
60
+ x = self.pool(x)
61
+
62
+ # [batch, hidden_size]
63
+ x = x.squeeze(-1)
64
+
65
+ # [batch, context_size * output_size]
66
+ out = self.fc(x)
67
+
68
+ # [batch, context_size, output_size]
69
+ out = out.view(
70
+ x.size(0),
71
+ self.context_size,
72
+ self.output_size,
73
+ )
74
+
75
+ return out