WEBVTT

00:00:00.000 --> 00:00:00.500 align:middle line:90%


00:00:00.500 --> 00:00:01.988 align:middle line:90%
[SQUEAKING]

00:00:01.988 --> 00:00:03.476 align:middle line:90%
[RUSTLING]

00:00:03.476 --> 00:00:05.460 align:middle line:90%
[CLICKING]

00:00:05.460 --> 00:00:12.920 align:middle line:90%


00:00:12.920 --> 00:00:14.520 align:middle line:90%
JEREMY BERNSTEIN: Hello.

00:00:14.520 --> 00:00:17.240 align:middle line:90%
OK, I'm just going to start.

00:00:17.240 --> 00:00:23.320 align:middle line:90%
So here's a question.

00:00:23.320 --> 00:00:25.960 align:middle line:84%
And I'm genuinely
asking, would you

00:00:25.960 --> 00:00:29.700 align:middle line:84%
rather to scale the width or the
depth of your neural network?

00:00:29.700 --> 00:00:32.920 align:middle line:90%


00:00:32.920 --> 00:00:34.540 align:middle line:90%
Does anyone have any thoughts?

00:00:34.540 --> 00:00:38.040 align:middle line:90%


00:00:38.040 --> 00:00:40.480 align:middle line:90%
Wait, someone said something.

00:00:40.480 --> 00:00:43.600 align:middle line:90%
I heard someone say depth.

00:00:43.600 --> 00:00:44.400 align:middle line:90%
AUDIENCE: Depth.

00:00:44.400 --> 00:00:45.347 align:middle line:90%
JEREMY BERNSTEIN: Why?

00:00:45.347 --> 00:00:46.180 align:middle line:90%
AUDIENCE: No reason.

00:00:46.180 --> 00:00:47.720 align:middle line:90%
I know someone said it earlier.

00:00:47.720 --> 00:00:50.012 align:middle line:84%
JEREMY BERNSTEIN: Because
you've heard "deep learning,"

00:00:50.012 --> 00:00:51.560 align:middle line:90%
so it sounds-- not wider, yeah.

00:00:51.560 --> 00:00:53.920 align:middle line:90%
[LAUGHTER]

00:00:53.920 --> 00:00:54.420 align:middle line:90%
OK.

00:00:54.420 --> 00:00:56.320 align:middle line:90%
But it's not really clear.

00:00:56.320 --> 00:01:00.630 align:middle line:84%
How do you even answer
a question like that?

00:01:00.630 --> 00:01:01.830 align:middle line:90%
I don't know.

00:01:01.830 --> 00:01:06.510 align:middle line:84%
It's something we
got to think about.

00:01:06.510 --> 00:01:14.430 align:middle line:84%
So in today's lecture, we're
going to ask this question of,

00:01:14.430 --> 00:01:17.670 align:middle line:84%
what class of functions can
a neural network express?

00:01:17.670 --> 00:01:19.710 align:middle line:84%
And you may have
heard this statement

00:01:19.710 --> 00:01:22.150 align:middle line:84%
that neural networks
are universal function

00:01:22.150 --> 00:01:23.750 align:middle line:90%
approximators.

00:01:23.750 --> 00:01:28.430 align:middle line:84%
So we're going to think
about what that means.

00:01:28.430 --> 00:01:32.950 align:middle line:84%
And if a neural network is a
universal function approximator

00:01:32.950 --> 00:01:34.990 align:middle line:84%
and it can approximate
any function that we're

00:01:34.990 --> 00:01:37.430 align:middle line:84%
interested in, then how
should we make decisions

00:01:37.430 --> 00:01:39.390 align:middle line:90%
about the architecture?

00:01:39.390 --> 00:01:42.270 align:middle line:84%
Does it matter how
many layers there are?

00:01:42.270 --> 00:01:46.550 align:middle line:84%
Maybe a few layers is enough
if you make them really wide.

00:01:46.550 --> 00:01:49.710 align:middle line:90%
Should you have lots of layers?

00:01:49.710 --> 00:01:52.090 align:middle line:84%
And ultimately, what do
you want to do in practice?

00:01:52.090 --> 00:01:54.310 align:middle line:84%
And does thinking
about neural networks

00:01:54.310 --> 00:01:57.310 align:middle line:84%
through this lens of universal
function approximation

00:01:57.310 --> 00:02:00.340 align:middle line:90%
actually help you in practice?

00:02:00.340 --> 00:02:03.980 align:middle line:84%
I'm not going to claim to
answer all of these questions,

00:02:03.980 --> 00:02:08.699 align:middle line:84%
but at least try to provide some
beginning of a framework, one

00:02:08.699 --> 00:02:10.440 align:middle line:84%
framework for
thinking about them.

00:02:10.440 --> 00:02:14.380 align:middle line:90%


00:02:14.380 --> 00:02:19.140 align:middle line:84%
So in machine learning,
you can think of it a bit

00:02:19.140 --> 00:02:20.200 align:middle line:90%
like we have a puzzle.

00:02:20.200 --> 00:02:22.080 align:middle line:84%
And there are different
pieces of the puzzle.

00:02:22.080 --> 00:02:23.538 align:middle line:84%
We're not going to
solve the puzzle

00:02:23.538 --> 00:02:27.820 align:middle line:84%
until we understand all of the
pieces and put them together.

00:02:27.820 --> 00:02:32.820 align:middle line:84%
And I broke it up here
into three pieces.

00:02:32.820 --> 00:02:38.882 align:middle line:84%
Let's see, so the
first is approximation.

00:02:38.882 --> 00:02:40.340 align:middle line:84%
This is the question
of, does there

00:02:40.340 --> 00:02:44.260 align:middle line:84%
exist a neural network
in my model family

00:02:44.260 --> 00:02:45.520 align:middle line:90%
that fits the training data?

00:02:45.520 --> 00:02:47.460 align:middle line:84%
So given my
architecture, can it even

00:02:47.460 --> 00:02:49.540 align:middle line:84%
represent the
function that I want,

00:02:49.540 --> 00:02:51.700 align:middle line:90%
that fits the training data?

00:02:51.700 --> 00:02:54.000 align:middle line:84%
The second question
is optimization.

00:02:54.000 --> 00:02:56.070 align:middle line:90%
So let's suppose it does exist.

00:02:56.070 --> 00:03:00.490 align:middle line:84%
Then that's the question
of, can I find it?

00:03:00.490 --> 00:03:02.270 align:middle line:90%
And third is generalization.

00:03:02.270 --> 00:03:03.390 align:middle line:90%
So suppose it exists.

00:03:03.390 --> 00:03:04.930 align:middle line:90%
Suppose I can find it.

00:03:04.930 --> 00:03:08.930 align:middle line:84%
Then it's like, does it
do well on unseen data?

00:03:08.930 --> 00:03:11.130 align:middle line:84%
So I really want
all of those things

00:03:11.130 --> 00:03:14.490 align:middle line:84%
to happen if I'm going to solve
my machine learning problem.

00:03:14.490 --> 00:03:16.570 align:middle line:84%
But in this lecture,
we're only going

00:03:16.570 --> 00:03:22.068 align:middle line:84%
to look at the first
question about approximation.

00:03:22.068 --> 00:03:24.110 align:middle line:84%
But just to remind you
that there is this puzzle.

00:03:24.110 --> 00:03:25.370 align:middle line:84%
We really want to
solve the whole puzzle

00:03:25.370 --> 00:03:26.830 align:middle line:90%
and put the pieces together.

00:03:26.830 --> 00:03:30.690 align:middle line:90%


00:03:30.690 --> 00:03:33.930 align:middle line:90%
So here is a motivating problem.

00:03:33.930 --> 00:03:36.610 align:middle line:84%
So imagine I have
two-dimensional data.

00:03:36.610 --> 00:03:39.290 align:middle line:84%
It's got two
coordinates, x1 and x2.

00:03:39.290 --> 00:03:45.010 align:middle line:84%
And there's two classes, the red
crosses and the green circles.

00:03:45.010 --> 00:03:48.210 align:middle line:84%
And I've come up with
my, think of this

00:03:48.210 --> 00:03:50.450 align:middle line:90%
as a one-layer neural net.

00:03:50.450 --> 00:03:53.790 align:middle line:84%
It's just the
w-transpose x plus b.

00:03:53.790 --> 00:03:56.080 align:middle line:90%
And then I take ReLU of that.

00:03:56.080 --> 00:03:58.460 align:middle line:84%
If that's my model,
can I fit this data?

00:03:58.460 --> 00:04:01.080 align:middle line:90%


00:04:01.080 --> 00:04:02.320 align:middle line:90%
What's your name, sir?

00:04:02.320 --> 00:04:02.982 align:middle line:90%
AUDIENCE: Matt.

00:04:02.982 --> 00:04:03.940 align:middle line:90%
JEREMY BERNSTEIN: Matt?

00:04:03.940 --> 00:04:04.280 align:middle line:90%
AUDIENCE: Yeah.

00:04:04.280 --> 00:04:05.780 align:middle line:84%
JEREMY BERNSTEIN:
Matt is saying no.

00:04:05.780 --> 00:04:06.960 align:middle line:90%
Why?

00:04:06.960 --> 00:04:10.560 align:middle line:84%
AUDIENCE: Don't you need two
linear separation planes?

00:04:10.560 --> 00:04:12.060 align:middle line:90%
JEREMY BERNSTEIN: Yeah.

00:04:12.060 --> 00:04:15.440 align:middle line:84%
Because this thing,
there's only one hyperplane

00:04:15.440 --> 00:04:18.120 align:middle line:90%
defined by w-transpose x.

00:04:18.120 --> 00:04:21.540 align:middle line:84%
And ReLU does not distort
the shape of the hyperplane.

00:04:21.540 --> 00:04:23.860 align:middle line:84%
It just applies
nonlinearity on the output.

00:04:23.860 --> 00:04:26.800 align:middle line:90%


00:04:26.800 --> 00:04:28.620 align:middle line:84%
This is basically
a linear separator.

00:04:28.620 --> 00:04:31.980 align:middle line:84%
And this data is not
linearly separable.

00:04:31.980 --> 00:04:36.800 align:middle line:84%
So this is supposed to be
recap from lecture one.

00:04:36.800 --> 00:04:38.860 align:middle line:90%
So hopefully, it's reasonable.

00:04:38.860 --> 00:04:42.120 align:middle line:84%
But we cannot linearly
separate this.

00:04:42.120 --> 00:04:44.680 align:middle line:84%
But what if I had a
two-layer neural net?

00:04:44.680 --> 00:04:47.560 align:middle line:84%
Do you think I could maybe
do it with two layers?

00:04:47.560 --> 00:04:48.720 align:middle line:90%
Potentially?

00:04:48.720 --> 00:04:49.240 align:middle line:90%
Yeah.

00:04:49.240 --> 00:04:49.540 align:middle line:90%
OK.

00:04:49.540 --> 00:04:50.940 align:middle line:84%
So I think we can do
it with two layers,

00:04:50.940 --> 00:04:52.398 align:middle line:84%
but you may need
to think about it.

00:04:52.398 --> 00:04:54.470 align:middle line:84%
But this is just supposed
to be a refresher.

00:04:54.470 --> 00:04:56.670 align:middle line:84%
But it's saying we've
got a kind of function

00:04:56.670 --> 00:04:59.272 align:middle line:84%
that we're trying to fit,
and we have a model family.

00:04:59.272 --> 00:05:01.230 align:middle line:84%
And we're asking, does
there exist the function

00:05:01.230 --> 00:05:04.470 align:middle line:84%
in that model family
that can fit that data?

00:05:04.470 --> 00:05:06.150 align:middle line:84%
So it's a kind of
warm-up for thinking

00:05:06.150 --> 00:05:12.030 align:middle line:84%
about this idea of
function approximation.

00:05:12.030 --> 00:05:14.430 align:middle line:84%
I have a second
motivating problem.

00:05:14.430 --> 00:05:19.350 align:middle line:84%
Does anyone recognize
this function?

00:05:19.350 --> 00:05:21.270 align:middle line:90%
You recognize it?

00:05:21.270 --> 00:05:22.210 align:middle line:90%
What is it?

00:05:22.210 --> 00:05:25.190 align:middle line:84%
AUDIENCE: Isn't it
Weierstrass' function?

00:05:25.190 --> 00:05:26.390 align:middle line:90%
JEREMY BERNSTEIN: Yeah.

00:05:26.390 --> 00:05:28.830 align:middle line:84%
It's Weierstrass'
function, which

00:05:28.830 --> 00:05:31.250 align:middle line:84%
is everywhere continuous
but nowhere differentiable.

00:05:31.250 --> 00:05:34.150 align:middle line:84%
It's a classic example of
a pathological function.

00:05:34.150 --> 00:05:36.870 align:middle line:84%
And the question I
wanted to ask is, do you

00:05:36.870 --> 00:05:39.530 align:middle line:84%
think you could fit this
function with a neural network?

00:05:39.530 --> 00:05:45.470 align:middle line:90%


00:05:45.470 --> 00:05:47.590 align:middle line:90%
Does anyone have a thought?

00:05:47.590 --> 00:05:48.190 align:middle line:90%
Matt?

00:05:48.190 --> 00:05:51.540 align:middle line:84%
AUDIENCE: Wouldn't you
have to go to infinity?

00:05:51.540 --> 00:05:53.420 align:middle line:84%
It would be like
asymptotically, I guess?

00:05:53.420 --> 00:05:54.420 align:middle line:84%
JEREMY BERNSTEIN: It
seems like you would need

00:05:54.420 --> 00:05:55.760 align:middle line:90%
some kind of asymptotic thing.

00:05:55.760 --> 00:05:58.560 align:middle line:84%
Yeah, honestly, I don't know
what the answer is to this.

00:05:58.560 --> 00:06:02.020 align:middle line:84%
I just thought it was
an interesting function.

00:06:02.020 --> 00:06:05.180 align:middle line:84%
And it's interesting to think,
can I fit it with a neural net?

00:06:05.180 --> 00:06:08.380 align:middle line:84%
I guess I suppose
the answer is, maybe.

00:06:08.380 --> 00:06:11.140 align:middle line:84%
And I thought a good final
project, if someone wants

00:06:11.140 --> 00:06:14.340 align:middle line:84%
an idea, is to take all
these pathological examples

00:06:14.340 --> 00:06:15.900 align:middle line:84%
that Weierstrass
was thinking about

00:06:15.900 --> 00:06:18.160 align:middle line:84%
and see if you can fit
them with a neural net.

00:06:18.160 --> 00:06:21.260 align:middle line:84%
But anyway, this is just more
motivation, just something

00:06:21.260 --> 00:06:23.300 align:middle line:90%
to think about.

00:06:23.300 --> 00:06:25.160 align:middle line:84%
OK, so this is a
fractal function.

00:06:25.160 --> 00:06:26.880 align:middle line:90%
It's defined recursively.

00:06:26.880 --> 00:06:28.100 align:middle line:90%
It's self-similar.

00:06:28.100 --> 00:06:30.340 align:middle line:84%
If you zoom into any
piece of the function,

00:06:30.340 --> 00:06:32.160 align:middle line:90%
it resembles the whole function.

00:06:32.160 --> 00:06:35.380 align:middle line:90%


00:06:35.380 --> 00:06:39.460 align:middle line:84%
So we want to now formalize
this approximation problem

00:06:39.460 --> 00:06:42.540 align:middle line:90%
that we've been talking about.

00:06:42.540 --> 00:06:47.300 align:middle line:90%
So given a family of curves G--

00:06:47.300 --> 00:06:49.490 align:middle line:84%
we think of these as
the curves that we

00:06:49.490 --> 00:06:52.410 align:middle line:84%
want to be able to
approximate ideally.

00:06:52.410 --> 00:06:54.910 align:middle line:84%
And we can choose
what we want G to be.

00:06:54.910 --> 00:06:58.670 align:middle line:84%
We could choose it to rule
out Weierstrass' function.

00:06:58.670 --> 00:07:00.590 align:middle line:84%
We could ask that G
is differentiable.

00:07:00.590 --> 00:07:04.730 align:middle line:84%
Or we're free to pick
whatever family of curves

00:07:04.730 --> 00:07:05.950 align:middle line:90%
we're interested in.

00:07:05.950 --> 00:07:08.050 align:middle line:90%
That's up to us.

00:07:08.050 --> 00:07:10.490 align:middle line:84%
And then we think of a
family of neural networks

00:07:10.490 --> 00:07:14.030 align:middle line:84%
F. This is like,
I'm going to say,

00:07:14.030 --> 00:07:17.410 align:middle line:84%
I care about a five-layer
MLP, a five-layer multilayer

00:07:17.410 --> 00:07:19.370 align:middle line:90%
perceptron.

00:07:19.370 --> 00:07:20.670 align:middle line:90%
That could be my family.

00:07:20.670 --> 00:07:22.650 align:middle line:84%
So in other words, a
neural architecture

00:07:22.650 --> 00:07:26.050 align:middle line:84%
could specify a
family of functions.

00:07:26.050 --> 00:07:28.730 align:middle line:84%
And then I want to say,
if I pick any curve

00:07:28.730 --> 00:07:32.330 align:middle line:84%
in my space of functions
or set of functions

00:07:32.330 --> 00:07:34.250 align:middle line:84%
that I'm interested
in, does there

00:07:34.250 --> 00:07:38.570 align:middle line:84%
exist a neural network in
my family of neural networks

00:07:38.570 --> 00:07:44.490 align:middle line:84%
that can fit that function
to some small error?

00:07:44.490 --> 00:07:48.080 align:middle line:84%
And we can come up with
different notions of error.

00:07:48.080 --> 00:07:50.800 align:middle line:84%
And OK, here, epsilon, you
should think of epsilon

00:07:50.800 --> 00:07:52.280 align:middle line:90%
as a small number.

00:07:52.280 --> 00:07:55.660 align:middle line:84%
So we choose
epsilon, and we ask,

00:07:55.660 --> 00:07:59.960 align:middle line:84%
can we approximate any function
in the family of curves G

00:07:59.960 --> 00:08:02.880 align:middle line:90%
to less than some small error?

00:08:02.880 --> 00:08:05.760 align:middle line:84%
And we have to ask, what
do we mean by the error?

00:08:05.760 --> 00:08:09.440 align:middle line:84%
So we could pick
whatever metric we want.

00:08:09.440 --> 00:08:14.100 align:middle line:84%
So one example, we could
call the L-infinity error,

00:08:14.100 --> 00:08:19.360 align:middle line:84%
which would be the max over
inputs of f of x minus g of x,

00:08:19.360 --> 00:08:20.940 align:middle line:84%
and then take the
absolute value.

00:08:20.940 --> 00:08:26.080 align:middle line:84%
So this, think of it
like L-infinity norm,

00:08:26.080 --> 00:08:27.960 align:middle line:90%
but for functions.

00:08:27.960 --> 00:08:30.680 align:middle line:84%
So it's the max discrepancy
between the two functions

00:08:30.680 --> 00:08:33.400 align:middle line:90%
anywhere on the input space.

00:08:33.400 --> 00:08:38.720 align:middle line:84%
But another equally
valid error measure, we

00:08:38.720 --> 00:08:41.539 align:middle line:84%
could call this the
L infinity measure.

00:08:41.539 --> 00:08:44.920 align:middle line:84%
Another measure would
be the L1 measure,

00:08:44.920 --> 00:08:47.190 align:middle line:84%
which would be the
integral over the input

00:08:47.190 --> 00:08:51.110 align:middle line:84%
space of the absolute value
of f of x minus g of x.

00:08:51.110 --> 00:08:52.590 align:middle line:84%
So that would be
another valid way

00:08:52.590 --> 00:08:54.673 align:middle line:84%
to measure discrepancy,
which is, I integrate over

00:08:54.673 --> 00:08:58.910 align:middle line:84%
the function and just sum
up all the differences

00:08:58.910 --> 00:09:00.210 align:middle line:90%
along the whole axis.

00:09:00.210 --> 00:09:03.510 align:middle line:90%


00:09:03.510 --> 00:09:06.270 align:middle line:84%
So in this lecture,
just for the purposes

00:09:06.270 --> 00:09:08.730 align:middle line:84%
of having something nice
to present in the lecture,

00:09:08.730 --> 00:09:12.310 align:middle line:84%
we're going to pick a
particular family of curves G,

00:09:12.310 --> 00:09:16.150 align:middle line:84%
that we're interested in
whether neural nets can

00:09:16.150 --> 00:09:17.670 align:middle line:90%
approximate this family.

00:09:17.670 --> 00:09:21.470 align:middle line:84%
And these are the Lipschitz
continuous functions.

00:09:21.470 --> 00:09:25.790 align:middle line:84%
So has anyone--
well, I think it's

00:09:25.790 --> 00:09:28.390 align:middle line:84%
nice to teach you about
Lipschitz functions regardless

00:09:28.390 --> 00:09:29.810 align:middle line:90%
of neural net approximation.

00:09:29.810 --> 00:09:32.050 align:middle line:84%
So that's another nice reason
to tell you about this.

00:09:32.050 --> 00:09:34.930 align:middle line:84%
So let's just think what a
Lipschitz continuous function

00:09:34.930 --> 00:09:35.430 align:middle line:90%
is.

00:09:35.430 --> 00:09:39.093 align:middle line:84%
It's a particular
notion of continuity.

00:09:39.093 --> 00:09:40.510 align:middle line:84%
So what we're going
to do is we're

00:09:40.510 --> 00:09:48.260 align:middle line:84%
going to say that a
function g from real numbers

00:09:48.260 --> 00:09:51.040 align:middle line:90%
to real numbers is L-Lipschitz.

00:09:51.040 --> 00:09:52.380 align:middle line:90%
So L is a number.

00:09:52.380 --> 00:09:53.920 align:middle line:90%
It's like, 10 or 5.

00:09:53.920 --> 00:09:57.220 align:middle line:84%
But it just measures how
Lipschitz the function is.

00:09:57.220 --> 00:10:04.060 align:middle line:84%
We'll say it's L-Lipschitz
if for any input point x,

00:10:04.060 --> 00:10:08.740 align:middle line:84%
the absolute value of g of x
plus delta x minus g of x--

00:10:08.740 --> 00:10:11.560 align:middle line:84%
so think, if I change
the inputs by a delta x,

00:10:11.560 --> 00:10:13.420 align:middle line:84%
I'm going to change
the function.

00:10:13.420 --> 00:10:15.900 align:middle line:84%
And the amount that
the function changes

00:10:15.900 --> 00:10:18.860 align:middle line:84%
is bounded by the
Lipschitz constant times

00:10:18.860 --> 00:10:20.900 align:middle line:90%
the size of delta x.

00:10:20.900 --> 00:10:25.860 align:middle line:84%
So this we call Lipschitz
continuity of g.

00:10:25.860 --> 00:10:30.100 align:middle line:84%
And there's an intuitive way to
think about what this is saying.

00:10:30.100 --> 00:10:33.140 align:middle line:84%
Imagine taking delta
x, take the limit

00:10:33.140 --> 00:10:35.820 align:middle line:84%
that delta x becomes
really small.

00:10:35.820 --> 00:10:39.160 align:middle line:84%
Can anyone recognize what
this looks a little bit like,

00:10:39.160 --> 00:10:41.560 align:middle line:84%
if we take the limit that
delta x becomes small?

00:10:41.560 --> 00:10:45.063 align:middle line:90%


00:10:45.063 --> 00:10:46.730 align:middle line:84%
AUDIENCE: Is it that
the slope is always

00:10:46.730 --> 00:10:48.327 align:middle line:90%
less than some constant L?

00:10:48.327 --> 00:10:49.410 align:middle line:90%
JEREMY BERNSTEIN: Exactly.

00:10:49.410 --> 00:10:50.550 align:middle line:90%
Yes, exactly.

00:10:50.550 --> 00:10:53.170 align:middle line:84%
So you should recognize that
if I take the limit that delta

00:10:53.170 --> 00:10:57.130 align:middle line:84%
x goes to 0, it's a bit like
saying that the derivative of g

00:10:57.130 --> 00:11:00.450 align:middle line:84%
is bounded by L. You
should see a definition

00:11:00.450 --> 00:11:03.130 align:middle line:84%
of the derivative
hiding in here.

00:11:03.130 --> 00:11:05.130 align:middle line:84%
So Lipschitz
continuity is a kind

00:11:05.130 --> 00:11:08.770 align:middle line:84%
of generalization of the notion
of having a bounded derivative.

00:11:08.770 --> 00:11:10.270 align:middle line:84%
Well, it may be
equivalent actually,

00:11:10.270 --> 00:11:11.645 align:middle line:84%
but you need to
think about that.

00:11:11.645 --> 00:11:13.410 align:middle line:84%
And just to draw
a picture, let's

00:11:13.410 --> 00:11:17.010 align:middle line:90%
suppose that we set x to 0.

00:11:17.010 --> 00:11:21.730 align:middle line:84%
So we think about
this x as being 0.

00:11:21.730 --> 00:11:26.450 align:middle line:84%
And then we ask how large
can g of delta x be.

00:11:26.450 --> 00:11:29.290 align:middle line:84%
And what this definition says,
if the function is L-Lipschitz,

00:11:29.290 --> 00:11:34.950 align:middle line:84%
we think about drawing
the line like y equals Lx.

00:11:34.950 --> 00:11:41.440 align:middle line:84%
And similarly, we draw the
other line y equals minus Lx.

00:11:41.440 --> 00:11:44.640 align:middle line:84%
And this definition implies
that whatever function

00:11:44.640 --> 00:11:47.100 align:middle line:84%
we have-- oops, it's not
supposed to be that squiggly.

00:11:47.100 --> 00:11:51.600 align:middle line:84%
It should be reasonably
not too-- anyway.

00:11:51.600 --> 00:11:54.900 align:middle line:84%
The function had better
lie within those bounds.

00:11:54.900 --> 00:11:58.100 align:middle line:84%
So the implication
of the definition,

00:11:58.100 --> 00:11:59.900 align:middle line:84%
if the function goes
through the origin,

00:11:59.900 --> 00:12:02.940 align:middle line:84%
it has to lie
within those bounds.

00:12:02.940 --> 00:12:05.140 align:middle line:84%
But the same argument
applies at any point.

00:12:05.140 --> 00:12:08.660 align:middle line:84%
So it places a cone
kind of shape maybe,

00:12:08.660 --> 00:12:10.080 align:middle line:90%
or a kind of bow tie.

00:12:10.080 --> 00:12:12.440 align:middle line:84%
And the function has to
belong to that bow tie

00:12:12.440 --> 00:12:15.240 align:middle line:84%
at any point along the
function, if that makes sense.

00:12:15.240 --> 00:12:19.000 align:middle line:84%
You imagine translating
the bow tie along.

00:12:19.000 --> 00:12:20.700 align:middle line:90%
So that's Lipschitz continuous.

00:12:20.700 --> 00:12:23.600 align:middle line:84%
And now I thought, well, another
great thing to introduce you to

00:12:23.600 --> 00:12:27.920 align:middle line:84%
is how to generalize
that notion to functions

00:12:27.920 --> 00:12:29.700 align:middle line:90%
with multiple inputs.

00:12:29.700 --> 00:12:32.600 align:middle line:84%
So now we're going to
think about generalizing

00:12:32.600 --> 00:12:38.480 align:middle line:84%
this Lipschitz notion
for g going from Rd to R.

00:12:38.480 --> 00:12:43.030 align:middle line:84%
And basically the inputs
are now going to be in Rd,

00:12:43.030 --> 00:12:46.510 align:middle line:84%
and the delta x's are
going to be in Rd.

00:12:46.510 --> 00:12:48.970 align:middle line:84%
And you see we don't really
need to change anything here.

00:12:48.970 --> 00:12:50.530 align:middle line:84%
The output is still
one dimensional,

00:12:50.530 --> 00:12:51.990 align:middle line:90%
so that still makes sense.

00:12:51.990 --> 00:12:56.150 align:middle line:84%
But this doesn't
really make sense,

00:12:56.150 --> 00:12:58.610 align:middle line:84%
because you can't take the
absolute value of a vector.

00:12:58.610 --> 00:13:00.690 align:middle line:84%
Or maybe you can, but
that's not what we want.

00:13:00.690 --> 00:13:03.750 align:middle line:90%


00:13:03.750 --> 00:13:05.430 align:middle line:84%
What we're going to
do is we're going

00:13:05.430 --> 00:13:09.630 align:middle line:84%
to change the absolute value to
the norm of the vector delta x.

00:13:09.630 --> 00:13:12.070 align:middle line:84%
And I'm going to pick a
particular norm that I really

00:13:12.070 --> 00:13:15.190 align:middle line:90%
like called the RMS norm.

00:13:15.190 --> 00:13:16.750 align:middle line:84%
So now, do you see
that this is now

00:13:16.750 --> 00:13:18.830 align:middle line:84%
a valid definition,
as long as I define

00:13:18.830 --> 00:13:21.910 align:middle line:90%
what I mean by the RMS norm?

00:13:21.910 --> 00:13:24.202 align:middle line:84%
So the idea is that
because the x is a vectors,

00:13:24.202 --> 00:13:25.910 align:middle line:84%
we need to measure
the size of the vector

00:13:25.910 --> 00:13:27.830 align:middle line:90%
on the right-hand side.

00:13:27.830 --> 00:13:34.710 align:middle line:84%
So I'm going to define
the RMS norm of a vector x

00:13:34.710 --> 00:13:38.700 align:middle line:84%
to be the square root
of 1 over d, which

00:13:38.700 --> 00:13:41.540 align:middle line:84%
is the dimension of the
vector, times the sum

00:13:41.540 --> 00:13:47.740 align:middle line:84%
of the coordinates xi
squared, which we can also

00:13:47.740 --> 00:13:50.140 align:middle line:84%
think of as just being
1 over square root d

00:13:50.140 --> 00:13:54.540 align:middle line:90%
times the Euclidean norm of x.

00:13:54.540 --> 00:13:57.340 align:middle line:84%
So I just want to
introduce you to the idea

00:13:57.340 --> 00:13:59.560 align:middle line:84%
that if you have an
object like a vector,

00:13:59.560 --> 00:14:03.420 align:middle line:84%
there's many different ways
to talk about how large it is.

00:14:03.420 --> 00:14:05.660 align:middle line:84%
The Euclidean norm
is one example,

00:14:05.660 --> 00:14:09.460 align:middle line:84%
but the Euclidean norm is
kind of a dimensional object.

00:14:09.460 --> 00:14:11.700 align:middle line:84%
You can think about
the RMS norm as a kind

00:14:11.700 --> 00:14:13.880 align:middle line:84%
of non-dimensional analog
of the Euclidean norm.

00:14:13.880 --> 00:14:16.980 align:middle line:84%
Because if all the entries
of the vector are 1, then

00:14:16.980 --> 00:14:20.260 align:middle line:84%
the RMS norm is 1, whereas the
Euclidean norm would be like,

00:14:20.260 --> 00:14:21.340 align:middle line:90%
square root d.

00:14:21.340 --> 00:14:23.340 align:middle line:84%
But anyway, the
point of this slide

00:14:23.340 --> 00:14:25.420 align:middle line:84%
is just to generalize
Lipschitzness

00:14:25.420 --> 00:14:27.120 align:middle line:84%
to functions with
multiple inputs,

00:14:27.120 --> 00:14:29.660 align:middle line:84%
and to point out that there's
many different possible ways

00:14:29.660 --> 00:14:32.780 align:middle line:84%
to measure the
size of something.

00:14:32.780 --> 00:14:39.130 align:middle line:84%
OK, but why were we
even talking about this?

00:14:39.130 --> 00:14:40.970 align:middle line:84%
The point is that
in this lecture,

00:14:40.970 --> 00:14:43.130 align:middle line:84%
in the first part
of this lecture,

00:14:43.130 --> 00:14:48.930 align:middle line:84%
we are going to prove something
akin to a universal function

00:14:48.930 --> 00:14:52.170 align:middle line:90%
approximation theorem.

00:14:52.170 --> 00:14:55.330 align:middle line:84%
And the class of
functions g that we're

00:14:55.330 --> 00:14:59.690 align:middle line:84%
interested in
approximating is going

00:14:59.690 --> 00:15:02.290 align:middle line:84%
to be the class of
L-Lipschitz functions

00:15:02.290 --> 00:15:04.390 align:middle line:90%
that map from the hypercube.

00:15:04.390 --> 00:15:08.210 align:middle line:84%
The domain or the input
space is the hypercube

00:15:08.210 --> 00:15:09.850 align:middle line:90%
to the real numbers.

00:15:09.850 --> 00:15:11.310 align:middle line:90%
So this is the hypercube.

00:15:11.310 --> 00:15:12.450 align:middle line:90%
And these are the reals.

00:15:12.450 --> 00:15:15.330 align:middle line:84%
And we're only considering
L-Lipschitz functions, so

00:15:15.330 --> 00:15:18.410 align:middle line:84%
functions with
bounded derivatives.

00:15:18.410 --> 00:15:21.570 align:middle line:84%
And what we're going to say
is that if we pick an error

00:15:21.570 --> 00:15:24.770 align:middle line:84%
tolerance that we're interested
in, then there exists

00:15:24.770 --> 00:15:30.410 align:middle line:84%
a three-layer ReLU network with
a certain number of units such

00:15:30.410 --> 00:15:34.960 align:middle line:84%
that this integral
notion of error--

00:15:34.960 --> 00:15:37.172 align:middle line:84%
summing up all the little
pieces and adding up

00:15:37.172 --> 00:15:38.880 align:middle line:84%
all the absolute values
of the difference

00:15:38.880 --> 00:15:40.880 align:middle line:84%
between our neural
net and the function

00:15:40.880 --> 00:15:46.720 align:middle line:84%
g-- that integral is less
than apparently 2 epsilon.

00:15:46.720 --> 00:15:49.900 align:middle line:84%
Does anyone have any
questions about this?

00:15:49.900 --> 00:15:54.840 align:middle line:84%
We could dwell on
it for a moment.

00:15:54.840 --> 00:15:58.380 align:middle line:84%
This is our goal basically
for the next 20 minutes.

00:15:58.380 --> 00:16:01.320 align:middle line:84%
Hopefully, we're just going to
try to prove this statement.

00:16:01.320 --> 00:16:04.020 align:middle line:84%
Why should we be interested
in such a statement?

00:16:04.020 --> 00:16:05.340 align:middle line:90%
That's another question.

00:16:05.340 --> 00:16:09.440 align:middle line:90%


00:16:09.440 --> 00:16:10.460 align:middle line:90%
Yeah, go ahead.

00:16:10.460 --> 00:16:15.360 align:middle line:84%
AUDIENCE: The 0, 1, the
mapping from 0 to 1 definition,

00:16:15.360 --> 00:16:17.580 align:middle line:84%
is that taking only
values from 0 and 1?

00:16:17.580 --> 00:16:20.120 align:middle line:90%
Is that what that means?

00:16:20.120 --> 00:16:22.760 align:middle line:84%
JEREMY BERNSTEIN: The interval
0, 1, with square brackets

00:16:22.760 --> 00:16:27.040 align:middle line:84%
refers to the interval of the
real line just between 0 and 1.

00:16:27.040 --> 00:16:29.980 align:middle line:84%
So that's just picking out
an interval on the real line.

00:16:29.980 --> 00:16:32.000 align:middle line:84%
And then raising
it to the power d

00:16:32.000 --> 00:16:34.930 align:middle line:84%
means that in all the dimensions
of our d-dimensional space,

00:16:34.930 --> 00:16:36.210 align:middle line:90%
we pick out that interval.

00:16:36.210 --> 00:16:38.070 align:middle line:84%
So it's like picking
out a cube in three

00:16:38.070 --> 00:16:40.070 align:middle line:90%
dimensions or a hypercube.

00:16:40.070 --> 00:16:44.590 align:middle line:84%
It's just supposed to be an
abstract way of writing down

00:16:44.590 --> 00:16:48.110 align:middle line:84%
the hypercube in d-dimensional
space, where all the axes are

00:16:48.110 --> 00:16:51.550 align:middle line:90%
between 0 and 1.

00:16:51.550 --> 00:16:54.277 align:middle line:90%
AUDIENCE: What is N here?

00:16:54.277 --> 00:16:55.110 align:middle line:90%
JEREMY BERNSTEIN: N?

00:16:55.110 --> 00:16:58.550 align:middle line:90%


00:16:58.550 --> 00:17:03.240 align:middle line:84%
Capital N is my symbol
for the number of units.

00:17:03.240 --> 00:17:07.589 align:middle line:84%
By units, we mean the number
of neurons in the ReLU MLP.

00:17:07.589 --> 00:17:11.750 align:middle line:84%
So N is the number of neurons
in the three-layer network.

00:17:11.750 --> 00:17:12.944 align:middle line:90%
AUDIENCE: OK.

00:17:12.944 --> 00:17:14.569 align:middle line:84%
AUDIENCE: I have a
couple of questions.

00:17:14.569 --> 00:17:18.150 align:middle line:84%
So first, are those
neurons per layer

00:17:18.150 --> 00:17:20.435 align:middle line:84%
or are they the total
number in the whole network?

00:17:20.435 --> 00:17:21.810 align:middle line:84%
JEREMY BERNSTEIN:
It's the total.

00:17:21.810 --> 00:17:24.015 align:middle line:84%
But honestly, because
it's three layers,

00:17:24.015 --> 00:17:25.890 align:middle line:84%
for the purpose of
understanding the lecture,

00:17:25.890 --> 00:17:28.490 align:middle line:84%
you can just pretend that
the number 3 is the number 1.

00:17:28.490 --> 00:17:31.560 align:middle line:84%
And that kind of question
doesn't really matter so much.

00:17:31.560 --> 00:17:38.220 align:middle line:84%
But I really mean actually
the sum total of the neurons.

00:17:38.220 --> 00:17:41.620 align:middle line:84%
AUDIENCE: So a
second question is,

00:17:41.620 --> 00:17:44.000 align:middle line:84%
when you say three-layer
ReLU network,

00:17:44.000 --> 00:17:45.560 align:middle line:84%
what precisely do
you mean by that?

00:17:45.560 --> 00:17:48.780 align:middle line:84%
So are there three
value functions in it?

00:17:48.780 --> 00:17:49.620 align:middle line:90%
For example--

00:17:49.620 --> 00:17:51.453 align:middle line:84%
JEREMY BERNSTEIN: I
precisely mean the thing

00:17:51.453 --> 00:17:54.540 align:middle line:84%
with three weight
matrices and two value

00:17:54.540 --> 00:17:56.500 align:middle line:84%
functions that follow
the first weight matrix

00:17:56.500 --> 00:17:57.760 align:middle line:90%
and the second weight matrix.

00:17:57.760 --> 00:18:00.040 align:middle line:84%
And then the third weight
matrix doesn't have a ReLU.

00:18:00.040 --> 00:18:00.582 align:middle line:90%
AUDIENCE: OK.

00:18:00.582 --> 00:18:03.100 align:middle line:84%
And then the final
question I had is,

00:18:03.100 --> 00:18:05.488 align:middle line:84%
this theorem seems
very specific.

00:18:05.488 --> 00:18:07.780 align:middle line:84%
I was just wondering if you
could speak a little bit as

00:18:07.780 --> 00:18:10.383 align:middle line:84%
to what you brought up--
why do we care about this?

00:18:10.383 --> 00:18:11.300 align:middle line:90%
JEREMY BERNSTEIN: Yes.

00:18:11.300 --> 00:18:12.020 align:middle line:90%
AUDIENCE: [INAUDIBLE]

00:18:12.020 --> 00:18:14.040 align:middle line:84%
JEREMY BERNSTEIN: So
that's a great question.

00:18:14.040 --> 00:18:17.940 align:middle line:84%
So the reason is because the
way to prove this theorem

00:18:17.940 --> 00:18:19.500 align:middle line:90%
is not super involved.

00:18:19.500 --> 00:18:23.100 align:middle line:84%
So the point is to show you
an example of such a theorem.

00:18:23.100 --> 00:18:26.460 align:middle line:84%
And you can see the whole
proof and how it works.

00:18:26.460 --> 00:18:29.060 align:middle line:84%
And then I'll point you
to some other references

00:18:29.060 --> 00:18:30.890 align:middle line:84%
where they do
different calculations

00:18:30.890 --> 00:18:32.030 align:middle line:90%
or different things.

00:18:32.030 --> 00:18:33.670 align:middle line:84%
It's just to show
you an example.

00:18:33.670 --> 00:18:36.470 align:middle line:84%
And then we're really going to
think about, is this relevant?

00:18:36.470 --> 00:18:38.710 align:middle line:84%
I'm not claiming that this
is an important result.

00:18:38.710 --> 00:18:40.755 align:middle line:84%
I'm just saying it's
something you can prove.

00:18:40.755 --> 00:18:42.130 align:middle line:84%
And then we'll
think a bit about,

00:18:42.130 --> 00:18:44.190 align:middle line:84%
is this an important
result or not?

00:18:44.190 --> 00:18:47.170 align:middle line:90%


00:18:47.170 --> 00:18:47.855 align:middle line:90%
Yeah?

00:18:47.855 --> 00:18:49.730 align:middle line:84%
AUDIENCE: Does this
theorem apply to networks

00:18:49.730 --> 00:18:51.610 align:middle line:90%
that are more than one layer?

00:18:51.610 --> 00:18:54.930 align:middle line:84%
In our notes from
the first lecture,

00:18:54.930 --> 00:18:57.690 align:middle line:84%
there was a section about
universal approximation theorem

00:18:57.690 --> 00:18:59.770 align:middle line:84%
for a single value
layer, where you have

00:18:59.770 --> 00:19:02.950 align:middle line:90%
a given amount of value units.

00:19:02.950 --> 00:19:05.850 align:middle line:84%
And in that single layer, you
could approximate a function

00:19:05.850 --> 00:19:07.270 align:middle line:90%
to arbitrary precision.

00:19:07.270 --> 00:19:08.450 align:middle line:90%
So does that apply?

00:19:08.450 --> 00:19:11.530 align:middle line:84%
Does this theorem apply to
layers that are more than one?

00:19:11.530 --> 00:19:14.510 align:middle line:84%
JEREMY BERNSTEIN: This applies
to a three-layer ReLU network.

00:19:14.510 --> 00:19:17.430 align:middle line:84%
So this is not the most general
thing that you can prove.

00:19:17.430 --> 00:19:19.950 align:middle line:84%
The result that you're talking
about is a different result.

00:19:19.950 --> 00:19:22.090 align:middle line:84%
This is just
another result where

00:19:22.090 --> 00:19:24.690 align:middle line:84%
there's a proof, which
we can put into some

00:19:24.690 --> 00:19:25.950 align:middle line:90%
slides and show to you.

00:19:25.950 --> 00:19:28.240 align:middle line:84%
It's just to get you
thinking about how

00:19:28.240 --> 00:19:29.540 align:middle line:90%
such a result could look.

00:19:29.540 --> 00:19:30.820 align:middle line:90%
But there are other results.

00:19:30.820 --> 00:19:32.653 align:middle line:84%
There are probably more
interesting versions

00:19:32.653 --> 00:19:34.880 align:middle line:84%
of this result. This
is just a result,

00:19:34.880 --> 00:19:38.800 align:middle line:84%
but it's about
three-layer network.

00:19:38.800 --> 00:19:45.200 align:middle line:84%
AUDIENCE: So for this, is it
possible to shift the interval

00:19:45.200 --> 00:19:46.962 align:middle line:90%
or scale the interval?

00:19:46.962 --> 00:19:47.920 align:middle line:90%
JEREMY BERNSTEIN: Yeah.

00:19:47.920 --> 00:19:49.960 align:middle line:84%
Anytime you see the
hypercube, you should think,

00:19:49.960 --> 00:19:51.980 align:middle line:84%
I can probably
just rescale that.

00:19:51.980 --> 00:19:54.380 align:middle line:84%
And I'm going to change
some scaling numbers.

00:19:54.380 --> 00:19:57.302 align:middle line:84%
I'd probably change maybe what
the Lipschitz constant is.

00:19:57.302 --> 00:19:58.760 align:middle line:84%
Yeah, you should
be able to rescale

00:19:58.760 --> 00:20:02.680 align:middle line:84%
it to be any hypercube
that has arbitrary

00:20:02.680 --> 00:20:05.020 align:middle line:84%
sizes in different
dimensions, basically,

00:20:05.020 --> 00:20:08.800 align:middle line:84%
or hypercuboid or
something like that.

00:20:08.800 --> 00:20:09.500 align:middle line:90%
Yes.

00:20:09.500 --> 00:20:13.680 align:middle line:90%


00:20:13.680 --> 00:20:17.500 align:middle line:84%
AUDIENCE: Why do we want to
constrain everything to cube?

00:20:17.500 --> 00:20:20.088 align:middle line:84%
Does this mean that
everything should be finite?

00:20:20.088 --> 00:20:22.380 align:middle line:84%
JEREMY BERNSTEIN: Yeah, that's
an interesting question.

00:20:22.380 --> 00:20:25.440 align:middle line:90%


00:20:25.440 --> 00:20:29.310 align:middle line:84%
Is that a toy thing
about this theorem?

00:20:29.310 --> 00:20:31.910 align:middle line:84%
But actually, if you think
about in real deep learning,

00:20:31.910 --> 00:20:37.630 align:middle line:84%
you usually normalize your
inputs to be coordinate-wise,

00:20:37.630 --> 00:20:39.162 align:middle line:90%
about 1 in magnitude.

00:20:39.162 --> 00:20:40.870 align:middle line:84%
You could actually
train a neural network

00:20:40.870 --> 00:20:42.453 align:middle line:84%
on an interesting
problem and actually

00:20:42.453 --> 00:20:45.630 align:middle line:84%
project all the inputs to live
within a hypercube, or maybe--

00:20:45.630 --> 00:20:47.170 align:middle line:90%
a hypercube.

00:20:47.170 --> 00:20:47.850 align:middle line:90%
Yeah, yeah.

00:20:47.850 --> 00:20:51.430 align:middle line:84%
So that part is actually
a realistic assumption

00:20:51.430 --> 00:20:54.530 align:middle line:84%
that my data lives in a
hypercube, for many problems,

00:20:54.530 --> 00:20:57.270 align:middle line:90%
not for all problems.

00:20:57.270 --> 00:20:58.903 align:middle line:90%
Yeah, last question.

00:20:58.903 --> 00:21:01.570 align:middle line:84%
AUDIENCE: So the 1 over epsilon
to the d seems pretty important.

00:21:01.570 --> 00:21:04.492 align:middle line:84%
Do you know if you relax
the L-Lipschitz that

00:21:04.492 --> 00:21:06.450 align:middle line:84%
will take a different
definition of continuous,

00:21:06.450 --> 00:21:09.590 align:middle line:84%
do you still remain with
that as the best you can do?

00:21:09.590 --> 00:21:13.470 align:middle line:84%
Or can you do better
than that power law?

00:21:13.470 --> 00:21:14.870 align:middle line:84%
JEREMY BERNSTEIN:
Yeah, I imagine

00:21:14.870 --> 00:21:18.030 align:middle line:84%
with more knowledge about
the problem structure,

00:21:18.030 --> 00:21:20.570 align:middle line:90%
you can do much better.

00:21:20.570 --> 00:21:23.510 align:middle line:90%


00:21:23.510 --> 00:21:26.110 align:middle line:84%
I'm not trying to claim that
this is an interesting result.

00:21:26.110 --> 00:21:27.870 align:middle line:84%
And we'll talk about
the limitations.

00:21:27.870 --> 00:21:31.210 align:middle line:84%
And I think you
would hope to do--

00:21:31.210 --> 00:21:31.925 align:middle line:90%
yeah, just look.

00:21:31.925 --> 00:21:34.050 align:middle line:84%
It's saying that you need
N-- the number of neurons

00:21:34.050 --> 00:21:37.130 align:middle line:84%
you need, you take the
Lipschitz constant,

00:21:37.130 --> 00:21:39.930 align:middle line:84%
and you raise it to the power of
the dimension, which is really,

00:21:39.930 --> 00:21:41.570 align:middle line:84%
that could potentially
be massive.

00:21:41.570 --> 00:21:44.830 align:middle line:84%
And also the smaller the
error of tolerance you want,

00:21:44.830 --> 00:21:46.330 align:middle line:84%
the bigger the
number of neurons you

00:21:46.330 --> 00:21:49.550 align:middle line:84%
need in a way that depends
exponentially on the dimension.

00:21:49.550 --> 00:21:55.530 align:middle line:84%
So that's really bad
dimension dependence.

00:21:55.530 --> 00:22:00.963 align:middle line:84%
And we don't, in practice,
really want to do that ever.

00:22:00.963 --> 00:22:03.130 align:middle line:84%
Yeah, it's great to point
out that the result itself

00:22:03.130 --> 00:22:06.050 align:middle line:84%
is not even the best possible
result we might hope for.

00:22:06.050 --> 00:22:07.450 align:middle line:84%
Whether we can do
better probably

00:22:07.450 --> 00:22:09.210 align:middle line:84%
depends on structure
in the problem.

00:22:09.210 --> 00:22:11.210 align:middle line:84%
And I don't know even if
under these conditions,

00:22:11.210 --> 00:22:15.610 align:middle line:84%
maybe there's something
better, but I'm not sure.

00:22:15.610 --> 00:22:19.230 align:middle line:84%
OK, let's just plow ahead
now for a little bit.

00:22:19.230 --> 00:22:24.440 align:middle line:84%
So the strategy that we're going
to adopt for proving this result

00:22:24.440 --> 00:22:27.080 align:middle line:84%
is, first of all, we're going
to pretend that our inputs are

00:22:27.080 --> 00:22:27.860 align:middle line:90%
one dimensional.

00:22:27.860 --> 00:22:30.402 align:middle line:84%
We're just going to treat the
case of one-dimensional inputs,

00:22:30.402 --> 00:22:31.520 align:middle line:90%
one-dimensional outputs.

00:22:31.520 --> 00:22:34.380 align:middle line:84%
And we're going to forget about
ReLU networks for a little bit.

00:22:34.380 --> 00:22:37.880 align:middle line:84%
And we're going to think about
approximating our function just

00:22:37.880 --> 00:22:41.620 align:middle line:84%
with rectangles
that we build up.

00:22:41.620 --> 00:22:43.440 align:middle line:84%
And you can see that
you can approximate

00:22:43.440 --> 00:22:47.560 align:middle line:84%
a function by taking
these rectangles

00:22:47.560 --> 00:22:50.200 align:middle line:90%
and putting them like this.

00:22:50.200 --> 00:22:53.157 align:middle line:84%
And you can see that as you
make the width of each rectangle

00:22:53.157 --> 00:22:55.240 align:middle line:84%
smaller, you get a better
and better approximation

00:22:55.240 --> 00:22:57.440 align:middle line:90%
to your function.

00:22:57.440 --> 00:22:59.960 align:middle line:84%
And that's something we
can quantify quite easily.

00:22:59.960 --> 00:23:01.600 align:middle line:84%
The second step is
then we're going

00:23:01.600 --> 00:23:03.600 align:middle line:84%
to generalize that
rectangle construction

00:23:03.600 --> 00:23:07.760 align:middle line:84%
to multidimensional inputs,
still ignoring ReLU networks.

00:23:07.760 --> 00:23:10.440 align:middle line:84%
And then the third
step is we're going

00:23:10.440 --> 00:23:13.800 align:middle line:84%
to show that with a
two-layer ReLU network,

00:23:13.800 --> 00:23:17.580 align:middle line:84%
you can approximate a
hyperrectangle of that kind.

00:23:17.580 --> 00:23:19.880 align:middle line:84%
So then you can just use the
third layer of the network

00:23:19.880 --> 00:23:22.670 align:middle line:84%
to linearly combine
all of your rectangles

00:23:22.670 --> 00:23:24.950 align:middle line:84%
and you'll be able to
approximate your function.

00:23:24.950 --> 00:23:26.852 align:middle line:90%
So this is the strategy.

00:23:26.852 --> 00:23:28.810 align:middle line:84%
First approximate the
function with rectangles,

00:23:28.810 --> 00:23:31.390 align:middle line:84%
then approximate it
with hyperrectangles.

00:23:31.390 --> 00:23:34.390 align:middle line:84%
Then show that a
two-layer ReLU network can

00:23:34.390 --> 00:23:36.170 align:middle line:90%
approximate a hyperrectangle.

00:23:36.170 --> 00:23:38.950 align:middle line:90%


00:23:38.950 --> 00:23:42.570 align:middle line:84%
OK, let's just plow on,
in the interest of time.

00:23:42.570 --> 00:23:45.470 align:middle line:90%


00:23:45.470 --> 00:23:50.150 align:middle line:84%
So the first step
is to construct

00:23:50.150 --> 00:23:52.390 align:middle line:90%
these rectangular strips.

00:23:52.390 --> 00:23:55.710 align:middle line:84%
Each strip is centered
on a discrete grid point.

00:23:55.710 --> 00:23:58.910 align:middle line:84%
And you can think that this is
like building a function-- f

00:23:58.910 --> 00:24:03.210 align:middle line:84%
of x is the sum
over alpha i times,

00:24:03.210 --> 00:24:05.990 align:middle line:84%
let's call this
the i-th indicator.

00:24:05.990 --> 00:24:10.670 align:middle line:84%
What you should think is that
alpha i is measuring the height

00:24:10.670 --> 00:24:14.030 align:middle line:90%
of the i-th rectangle alpha i.

00:24:14.030 --> 00:24:17.910 align:middle line:84%
And indicator i is
the function, which

00:24:17.910 --> 00:24:22.860 align:middle line:84%
is 1, if the input is in this
interval, and 0 otherwise.

00:24:22.860 --> 00:24:28.180 align:middle line:84%
And then if I sum up a bunch of
these indicators of this form,

00:24:28.180 --> 00:24:31.160 align:middle line:84%
that corresponds to this
approximation to the function.

00:24:31.160 --> 00:24:33.760 align:middle line:90%


00:24:33.760 --> 00:24:35.320 align:middle line:90%
Is that clear?

00:24:35.320 --> 00:24:35.820 align:middle line:90%
OK.

00:24:35.820 --> 00:24:39.580 align:middle line:84%
So the next step is basically
what we're going to ask

00:24:39.580 --> 00:24:41.520 align:middle line:84%
is, if we approximate a
function in this way--

00:24:41.520 --> 00:24:44.540 align:middle line:84%
let's say we make the
left edge of the rectangle

00:24:44.540 --> 00:24:46.580 align:middle line:90%
actually touch the function.

00:24:46.580 --> 00:24:52.860 align:middle line:84%
We need to ask, how
big is this gap?

00:24:52.860 --> 00:24:54.380 align:middle line:90%
How big can it be?

00:24:54.380 --> 00:24:57.100 align:middle line:90%
And the observation is--

00:24:57.100 --> 00:24:59.380 align:middle line:84%
OK, does anyone
have an idea of how

00:24:59.380 --> 00:25:04.100 align:middle line:84%
to place a constraint on how
large that difference can be?

00:25:04.100 --> 00:25:04.600 align:middle line:90%
Yeah?

00:25:04.600 --> 00:25:06.308 align:middle line:84%
AUDIENCE: We use the
Lipschitz condition.

00:25:06.308 --> 00:25:08.297 align:middle line:84%
It would just be the
area of the triangle.

00:25:08.297 --> 00:25:09.380 align:middle line:90%
JEREMY BERNSTEIN: Exactly.

00:25:09.380 --> 00:25:11.373 align:middle line:84%
So the observation is
that we're assuming

00:25:11.373 --> 00:25:13.540 align:middle line:84%
that the function g that
we're trying to approximate

00:25:13.540 --> 00:25:14.900 align:middle line:90%
is L-Lipschitz.

00:25:14.900 --> 00:25:17.140 align:middle line:84%
Like we said, that places
a bound on how large

00:25:17.140 --> 00:25:18.510 align:middle line:90%
its derivative can be.

00:25:18.510 --> 00:25:21.690 align:middle line:84%
And that gives us a constraint
on how big this triangle can be.

00:25:21.690 --> 00:25:24.650 align:middle line:84%
Or the triangle could go
under, but it tells us

00:25:24.650 --> 00:25:26.630 align:middle line:90%
how big this thing can be.

00:25:26.630 --> 00:25:30.210 align:middle line:90%


00:25:30.210 --> 00:25:33.730 align:middle line:84%
And so we're going to ask,
given N strips, where N,

00:25:33.730 --> 00:25:37.190 align:middle line:84%
we can ramp it up or we can
make N as big as we want.

00:25:37.190 --> 00:25:40.050 align:middle line:84%
But given that we
have N of them,

00:25:40.050 --> 00:25:42.130 align:middle line:90%
what is the approximation error?

00:25:42.130 --> 00:25:45.530 align:middle line:84%
And the claim is that
the approximation error

00:25:45.530 --> 00:25:47.650 align:middle line:84%
is bounded by the
Lipschitz constant divided

00:25:47.650 --> 00:25:51.890 align:middle line:90%
by 2 times the number of strips.

00:25:51.890 --> 00:25:55.470 align:middle line:84%
And so importantly, if the
Lipschitz constant gets bigger,

00:25:55.470 --> 00:25:58.150 align:middle line:84%
then if I keep the number
of strips the same,

00:25:58.150 --> 00:25:59.710 align:middle line:90%
the error would get bigger.

00:25:59.710 --> 00:26:02.830 align:middle line:84%
But if I make the number
of rectangles larger,

00:26:02.830 --> 00:26:04.070 align:middle line:90%
the error would go down.

00:26:04.070 --> 00:26:10.170 align:middle line:90%


00:26:10.170 --> 00:26:13.370 align:middle line:84%
So to prove this, can you see
that basically, we just need

00:26:13.370 --> 00:26:16.360 align:middle line:84%
to do the geometry of what this
triangle can look like and then

00:26:16.360 --> 00:26:18.320 align:middle line:90%
sum it up over all the strips?

00:26:18.320 --> 00:26:20.960 align:middle line:90%
So let's just do that quickly.

00:26:20.960 --> 00:26:27.720 align:middle line:84%
So if we have N strips, each
rectangle is width 1 over N.

00:26:27.720 --> 00:26:31.080 align:middle line:84%
And the height of the
triangle, if you just

00:26:31.080 --> 00:26:33.340 align:middle line:84%
think through the
definition of Lipschitzness,

00:26:33.340 --> 00:26:36.400 align:middle line:84%
you just multiply the
width of the strip by L.

00:26:36.400 --> 00:26:39.600 align:middle line:84%
And that will give you the
maximum possible height.

00:26:39.600 --> 00:26:44.080 align:middle line:84%
And then, of course, then
the triangle has area--

00:26:44.080 --> 00:26:47.040 align:middle line:84%
What's the area
of this triangle?

00:26:47.040 --> 00:26:48.260 align:middle line:90%
Can someone tell me?

00:26:48.260 --> 00:26:49.440 align:middle line:90%
[LAUGHTER]

00:26:49.440 --> 00:26:51.180 align:middle line:90%
Oh, wait I remember.

00:26:51.180 --> 00:26:54.140 align:middle line:84%
It's that 1/2 L over
N squared, right?

00:26:54.140 --> 00:26:56.200 align:middle line:90%
The 1/2 base times height.

00:26:56.200 --> 00:26:58.700 align:middle line:84%
And this means that
the total error--

00:26:58.700 --> 00:27:02.760 align:middle line:84%
remember that we're dealing with
the L1 error of the function

00:27:02.760 --> 00:27:04.080 align:middle line:90%
approximation--

00:27:04.080 --> 00:27:09.920 align:middle line:84%
is just going to be like N times
the error of each rectangle.

00:27:09.920 --> 00:27:14.100 align:middle line:84%
And that corresponds to actually
the area of the triangle,

00:27:14.100 --> 00:27:15.970 align:middle line:84%
the maximum possible
area of the triangle,

00:27:15.970 --> 00:27:17.550 align:middle line:90%
which is 1/2 L over N squared.

00:27:17.550 --> 00:27:23.510 align:middle line:84%
So this is just
1/2 L over N. OK,

00:27:23.510 --> 00:27:25.330 align:middle line:84%
now we just need to
flip things around.

00:27:25.330 --> 00:27:28.050 align:middle line:84%
That's the largest possible
thing the error could be.

00:27:28.050 --> 00:27:30.010 align:middle line:84%
But what we wanted is a
statement of the form--

00:27:30.010 --> 00:27:35.150 align:middle line:84%
if we have this many rectangles,
then the error cannot exceed

00:27:35.150 --> 00:27:36.190 align:middle line:90%
epsilon.

00:27:36.190 --> 00:27:39.110 align:middle line:84%
So we need to flip
the statement around.

00:27:39.110 --> 00:27:43.710 align:middle line:84%
So that corresponds, the way you
do it mentally is just you say,

00:27:43.710 --> 00:27:47.550 align:middle line:90%
epsilon is 1/2 L over N.

00:27:47.550 --> 00:27:49.630 align:middle line:84%
So then to achieve
error epsilon, we just

00:27:49.630 --> 00:27:50.850 align:middle line:90%
rearrange this equation.

00:27:50.850 --> 00:27:56.190 align:middle line:84%
We get N needs to be greater
than or equal to 1/2 L over

00:27:56.190 --> 00:27:56.890 align:middle line:90%
epsilon.

00:27:56.890 --> 00:28:00.350 align:middle line:90%


00:28:00.350 --> 00:28:04.150 align:middle line:84%
And if I don't care
about this factor of 1/2,

00:28:04.150 --> 00:28:06.035 align:middle line:84%
I'm free to just also
erase that factor,

00:28:06.035 --> 00:28:07.410 align:middle line:84%
because it doesn't
really matter.

00:28:07.410 --> 00:28:11.270 align:middle line:90%


00:28:11.270 --> 00:28:14.820 align:middle line:84%
The form of this statement, is
if I set N greater than L over

00:28:14.820 --> 00:28:18.980 align:middle line:84%
epsilon, then I will have an
error epsilon less than 1/2 L

00:28:18.980 --> 00:28:22.300 align:middle line:84%
over N. So I'm just rearranging
the equation and being careful

00:28:22.300 --> 00:28:25.100 align:middle line:90%
about the inequalities.

00:28:25.100 --> 00:28:28.440 align:middle line:90%
So this is step one.

00:28:28.440 --> 00:28:30.020 align:middle line:84%
We know how many
rectangles we need

00:28:30.020 --> 00:28:32.780 align:middle line:84%
to get a certain error in
the one-dimensional case

00:28:32.780 --> 00:28:34.460 align:middle line:90%
with one-dimensional inputs.

00:28:34.460 --> 00:28:36.820 align:middle line:84%
Let's think about how
things change if we

00:28:36.820 --> 00:28:41.740 align:middle line:90%
have multidimensional inputs.

00:28:41.740 --> 00:28:49.420 align:middle line:84%
And the thing that changes is
we now have this hyperrectangle,

00:28:49.420 --> 00:28:50.898 align:middle line:90%
just think 3D.

00:28:50.898 --> 00:28:52.940 align:middle line:84%
Whenever you need to think
about high dimensions,

00:28:52.940 --> 00:28:56.520 align:middle line:84%
you just think about
things in three dimensions.

00:28:56.520 --> 00:29:00.940 align:middle line:90%


00:29:00.940 --> 00:29:04.900 align:middle line:84%
What I'm doing, I'm just
repeating the calculation,

00:29:04.900 --> 00:29:09.300 align:middle line:84%
the total error, and then I'm
breaking it up into pieces.

00:29:09.300 --> 00:29:13.470 align:middle line:84%
So this piece is just the
number of hyperrectangles.

00:29:13.470 --> 00:29:16.770 align:middle line:90%


00:29:16.770 --> 00:29:21.970 align:middle line:84%
This piece is the error
per hyperrectangle.

00:29:21.970 --> 00:29:24.970 align:middle line:84%
No, we think about the
hyperrectangle as having

00:29:24.970 --> 00:29:28.210 align:middle line:84%
a little cap on top, where
the hyperrectangle differs

00:29:28.210 --> 00:29:33.170 align:middle line:90%
from the surface G.

00:29:33.170 --> 00:29:37.150 align:middle line:84%
And we know the surface area of
the top of the hyperrectangle,

00:29:37.150 --> 00:29:40.130 align:middle line:84%
because it's just the width
of the hyperrectangle raised

00:29:40.130 --> 00:29:42.510 align:middle line:84%
to the number of
input dimensions.

00:29:42.510 --> 00:29:44.910 align:middle line:84%
So we know the surface area
of the top of the rectangle.

00:29:44.910 --> 00:29:45.713 align:middle line:90%
That's easy.

00:29:45.713 --> 00:29:47.130 align:middle line:84%
And then we just
use the Lipschitz

00:29:47.130 --> 00:29:49.330 align:middle line:84%
constant to measure the
maximum possible height

00:29:49.330 --> 00:29:50.615 align:middle line:90%
of the hyperrectangle.

00:29:50.615 --> 00:29:52.490 align:middle line:84%
So it's just repeating
exactly the same thing

00:29:52.490 --> 00:29:55.910 align:middle line:84%
that we did, except instead
of the error being a triangle,

00:29:55.910 --> 00:29:58.230 align:middle line:84%
it's some kind of
generalization of a triangle.

00:29:58.230 --> 00:30:01.090 align:middle line:84%
So I'll probably just leave
people to think that through

00:30:01.090 --> 00:30:03.790 align:middle line:84%
by looking at the slides, if you
want to think more about that.

00:30:03.790 --> 00:30:06.130 align:middle line:90%
But this is the height.

00:30:06.130 --> 00:30:07.830 align:middle line:90%
Let's call it the error cap.

00:30:07.830 --> 00:30:20.728 align:middle line:90%


00:30:20.728 --> 00:30:22.160 align:middle line:90%
Is this right?

00:30:22.160 --> 00:30:24.390 align:middle line:90%


00:30:24.390 --> 00:30:25.140 align:middle line:90%
Yeah, it is right.

00:30:25.140 --> 00:30:27.098 align:middle line:84%
And then if you think
about, what's the surface

00:30:27.098 --> 00:30:29.020 align:middle line:90%
area of the top of the cap?

00:30:29.020 --> 00:30:31.840 align:middle line:84%
Well, if there's
N hyperrectangles,

00:30:31.840 --> 00:30:34.740 align:middle line:84%
if I sum up all the surface
areas of all of them,

00:30:34.740 --> 00:30:38.080 align:middle line:90%
I should get 1.

00:30:38.080 --> 00:30:40.600 align:middle line:84%
Yeah, so then the
surface area of each one

00:30:40.600 --> 00:30:43.860 align:middle line:84%
should better be 1 over N. So
we'll just call this the area.

00:30:43.860 --> 00:30:47.000 align:middle line:84%
If anyone wants to think through
this more carefully, please

00:30:47.000 --> 00:30:49.220 align:middle line:90%
do so.

00:30:49.220 --> 00:30:52.150 align:middle line:84%
And the claim is that we've
just done the same calculation.

00:30:52.150 --> 00:30:54.400 align:middle line:84%
The only remaining thing is
to rearrange it in the way

00:30:54.400 --> 00:30:55.640 align:middle line:90%
that we just did.

00:30:55.640 --> 00:31:01.120 align:middle line:84%
And this would say that we
should need L over epsilon

00:31:01.120 --> 00:31:02.480 align:middle line:90%
to the d.

00:31:02.480 --> 00:31:05.800 align:middle line:84%
So I do the same
trick of setting

00:31:05.800 --> 00:31:10.310 align:middle line:84%
epsilon equal to this thing,
and then just rearranging.

00:31:10.310 --> 00:31:15.990 align:middle line:84%
Now, let's just zoom
out and just try

00:31:15.990 --> 00:31:18.470 align:middle line:90%
to assess where we're up to.

00:31:18.470 --> 00:31:20.270 align:middle line:84%
So what we're imagining
is that we now

00:31:20.270 --> 00:31:23.730 align:middle line:84%
have a function g with
multiple inputs and one output,

00:31:23.730 --> 00:31:27.230 align:middle line:84%
which we can think
of as a surface.

00:31:27.230 --> 00:31:29.510 align:middle line:84%
That would be the
case where there's

00:31:29.510 --> 00:31:31.210 align:middle line:90%
two inputs and one output.

00:31:31.210 --> 00:31:32.090 align:middle line:90%
It's like a surface.

00:31:32.090 --> 00:31:36.110 align:middle line:84%
And then in higher dimensions,
it's a generalization of that.

00:31:36.110 --> 00:31:40.070 align:middle line:84%
And the claim is that if
we want to approximate this

00:31:40.070 --> 00:31:42.750 align:middle line:84%
with N hyperrectangles
and we want

00:31:42.750 --> 00:31:45.770 align:middle line:84%
to get a certain notion of
the integral of the error,

00:31:45.770 --> 00:31:48.950 align:middle line:84%
that the number that we
need to get an error epsilon

00:31:48.950 --> 00:31:51.110 align:middle line:90%
is L over epsilon to the d.

00:31:51.110 --> 00:31:53.750 align:middle line:84%
And this should remind you a
lot of the thing which appeared

00:31:53.750 --> 00:31:55.670 align:middle line:90%
in the theorem statement.

00:31:55.670 --> 00:31:58.030 align:middle line:84%
And the reason for that
is that the next step

00:31:58.030 --> 00:31:59.510 align:middle line:84%
is to show that
we can approximate

00:31:59.510 --> 00:32:04.470 align:middle line:84%
the rectangle with a ReLU,
with a two-layer ReLU network.

00:32:04.470 --> 00:32:06.780 align:middle line:84%
So the number of
small two-layer ReLU

00:32:06.780 --> 00:32:09.575 align:middle line:84%
networks we're going to need
is L over epsilon to the d.

00:32:09.575 --> 00:32:11.700 align:middle line:84%
And then we just need to
count how many neurons are

00:32:11.700 --> 00:32:13.760 align:middle line:90%
in each of those ReLU networks.

00:32:13.760 --> 00:32:16.740 align:middle line:90%


00:32:16.740 --> 00:32:24.660 align:middle line:84%
So I just decided to
just zoom over this.

00:32:24.660 --> 00:32:26.260 align:middle line:84%
Because I want to
show you actually

00:32:26.260 --> 00:32:27.780 align:middle line:90%
how to construct this thing.

00:32:27.780 --> 00:32:32.580 align:middle line:84%
But basically the claim is
that there's a one-layer--

00:32:32.580 --> 00:32:33.220 align:middle line:90%
is this one?

00:32:33.220 --> 00:32:35.140 align:middle line:84%
OK, we would call
this a two-layer ReLU

00:32:35.140 --> 00:32:41.540 align:middle line:84%
network with a parameter
c, a kind of weight c.

00:32:41.540 --> 00:32:43.940 align:middle line:90%
And if we take c to infinity--

00:32:43.940 --> 00:32:46.100 align:middle line:84%
which seems like a bit of
a trick, but OK, that's

00:32:46.100 --> 00:32:47.740 align:middle line:90%
what we're going to do--

00:32:47.740 --> 00:32:50.140 align:middle line:84%
the claim is that
that network converges

00:32:50.140 --> 00:32:55.340 align:middle line:84%
to this function, which is
exactly this single rectangle

00:32:55.340 --> 00:32:56.580 align:middle line:90%
in one dimension.

00:32:56.580 --> 00:32:59.700 align:middle line:84%
So I just want to
take a bit of a risk

00:32:59.700 --> 00:33:05.930 align:middle line:84%
by switching off and hoping that
I'm able to reshare my iPad.

00:33:05.930 --> 00:33:12.730 align:middle line:84%
I'll just try to build
this thing for you.

00:33:12.730 --> 00:33:15.730 align:middle line:90%
So this is a ReLU.

00:33:15.730 --> 00:33:17.450 align:middle line:84%
By the way, this
kind of manipulation

00:33:17.450 --> 00:33:20.850 align:middle line:84%
is quite helpful on some of the
homework problems, probably.

00:33:20.850 --> 00:33:21.950 align:middle line:90%
This, I really recommend.

00:33:21.950 --> 00:33:25.030 align:middle line:84%
I'm not being paid by
this graphing software,

00:33:25.030 --> 00:33:27.690 align:middle line:90%
but it's really great.

00:33:27.690 --> 00:33:32.130 align:middle line:84%
And what I want to show you
is, so here I've got ReLU of x.

00:33:32.130 --> 00:33:35.270 align:middle line:84%
And then I'm just going to
subtract ReLU of x minus 1.

00:33:35.270 --> 00:33:39.130 align:middle line:90%
And you see it flattens it out.

00:33:39.130 --> 00:33:42.470 align:middle line:84%
And then say I want
to curve it down,

00:33:42.470 --> 00:33:44.490 align:middle line:90%
so I subtract another ReLU.

00:33:44.490 --> 00:33:48.410 align:middle line:84%
And then I want to flatten off
that little flat bit that's

00:33:48.410 --> 00:33:49.030 align:middle line:90%
not flat.

00:33:49.030 --> 00:33:50.590 align:middle line:90%
So I'm just going to add a ReLU.

00:33:50.590 --> 00:33:53.710 align:middle line:84%
And each time, I'm
translating them one along.

00:33:53.710 --> 00:33:56.250 align:middle line:90%


00:33:56.250 --> 00:33:58.710 align:middle line:84%
But now I'm like, hey, that
doesn't look like a rectangle.

00:33:58.710 --> 00:34:00.890 align:middle line:90%
The slopes are too non-slope--

00:34:00.890 --> 00:34:02.990 align:middle line:84%
I want them to be
slopier or something.

00:34:02.990 --> 00:34:06.320 align:middle line:90%


00:34:06.320 --> 00:34:10.300 align:middle line:84%
OK, I'm going to insert a
constant c, which is initially 1

00:34:10.300 --> 00:34:12.199 align:middle line:90%
so it's not changing anything.

00:34:12.199 --> 00:34:14.600 align:middle line:84%
And I'm just going to use
this constant to increase

00:34:14.600 --> 00:34:18.760 align:middle line:90%
the slope of all the curves.

00:34:18.760 --> 00:34:21.315 align:middle line:84%
And let's put it
between 1 and 10.

00:34:21.315 --> 00:34:22.940 align:middle line:84%
And then I'm just
going to increase it.

00:34:22.940 --> 00:34:26.560 align:middle line:84%
And I'm like, oh, it's
making the slopes slopier,

00:34:26.560 --> 00:34:31.719 align:middle line:84%
but it's also squeezing
them together.

00:34:31.719 --> 00:34:38.840 align:middle line:84%
And so I can solve that problem
just by translating them back,

00:34:38.840 --> 00:34:39.920 align:middle line:90%
I hope.

00:34:39.920 --> 00:34:41.699 align:middle line:84%
That doesn't look
like what I wanted.

00:34:41.699 --> 00:34:45.679 align:middle line:90%


00:34:45.679 --> 00:34:47.242 align:middle line:90%
Let's see.

00:34:47.242 --> 00:34:48.659 align:middle line:84%
Oh, now it looks
like what I want.

00:34:48.659 --> 00:34:51.600 align:middle line:90%
So I just translated them back.

00:34:51.600 --> 00:34:54.659 align:middle line:84%
And then I'm just going to
let c get a little bigger,

00:34:54.659 --> 00:34:57.220 align:middle line:84%
maybe 1,000, and
make it a bit bigger.

00:34:57.220 --> 00:34:58.680 align:middle line:90%
And you see, OK, I did it.

00:34:58.680 --> 00:35:01.000 align:middle line:90%
That's just combining ReLUs.

00:35:01.000 --> 00:35:03.270 align:middle line:84%
That's technically a
two-layer neural net.

00:35:03.270 --> 00:35:06.230 align:middle line:84%
And I could approximate
a square function.

00:35:06.230 --> 00:35:10.990 align:middle line:84%
So that's what this slide
is doing, which I now

00:35:10.990 --> 00:35:12.330 align:middle line:90%
hope I can bring back.

00:35:12.330 --> 00:35:16.390 align:middle line:90%


00:35:16.390 --> 00:35:19.947 align:middle line:84%
So that is the demo
of approximating.

00:35:19.947 --> 00:35:22.030 align:middle line:84%
And you can see that using
that graphing software,

00:35:22.030 --> 00:35:23.470 align:middle line:84%
you can just play
around with things

00:35:23.470 --> 00:35:24.595 align:middle line:90%
that you can figure it out.

00:35:24.595 --> 00:35:25.130 align:middle line:90%
Go ahead.

00:35:25.130 --> 00:35:26.505 align:middle line:84%
AUDIENCE: What's
the relationship

00:35:26.505 --> 00:35:30.167 align:middle line:84%
between a two-layer ReLU
network and a Riemann sum?

00:35:30.167 --> 00:35:31.750 align:middle line:84%
JEREMY BERNSTEIN:
Yeah, so we're going

00:35:31.750 --> 00:35:35.790 align:middle line:84%
to use one copy of
this two-layer network

00:35:35.790 --> 00:35:38.750 align:middle line:84%
to approximate each
rectangle, to approximate one

00:35:38.750 --> 00:35:42.230 align:middle line:84%
of the rectangles in that
Riemann sum kind of thing.

00:35:42.230 --> 00:35:43.990 align:middle line:84%
And then we're going
to weight each.

00:35:43.990 --> 00:35:47.910 align:middle line:84%
And you have to insert factors
to translate it to the position

00:35:47.910 --> 00:35:48.530 align:middle line:90%
that you want.

00:35:48.530 --> 00:35:50.590 align:middle line:84%
So you translate it,
and then you multiply it

00:35:50.590 --> 00:35:54.290 align:middle line:84%
by a little scalar to get it
to be the size that you want.

00:35:54.290 --> 00:35:58.230 align:middle line:84%
And that scalar would
correspond to the third layer

00:35:58.230 --> 00:36:01.860 align:middle line:84%
in the network, to set the size
of the rectangle, how tall it is

00:36:01.860 --> 00:36:03.540 align:middle line:90%
or what its height is.

00:36:03.540 --> 00:36:06.026 align:middle line:90%
Yeah?

00:36:06.026 --> 00:36:07.780 align:middle line:84%
AUDIENCE: Why do we
restrict ourselves

00:36:07.780 --> 00:36:09.560 align:middle line:90%
to L-Lipschitz functions?

00:36:09.560 --> 00:36:13.887 align:middle line:84%
Wouldn't that be for
any kind of function?

00:36:13.887 --> 00:36:16.220 align:middle line:84%
JEREMY BERNSTEIN: Yeah, it's
to get a sense of the error

00:36:16.220 --> 00:36:17.520 align:middle line:90%
when things are finite.

00:36:17.520 --> 00:36:18.620 align:middle line:84%
Because we're not
interested in actually

00:36:18.620 --> 00:36:21.120 align:middle line:84%
taking the limit that the number
of strips goes to infinity.

00:36:21.120 --> 00:36:23.280 align:middle line:84%
We want to know for a
finite number of strips.

00:36:23.280 --> 00:36:26.160 align:middle line:84%
So at least that's how we're
using the Lipschitzness.

00:36:26.160 --> 00:36:29.880 align:middle line:84%
But, yeah, I think
it's necessary.

00:36:29.880 --> 00:36:31.160 align:middle line:90%
I'm not 100% sure.

00:36:31.160 --> 00:36:33.740 align:middle line:84%
But we can maybe talk about
that after the lecture.

00:36:33.740 --> 00:36:38.540 align:middle line:84%
Let me just keep going, because
I'm a bit worried about time.

00:36:38.540 --> 00:36:39.840 align:middle line:90%
This is just half the lecture.

00:36:39.840 --> 00:36:42.048 align:middle line:84%
And I guess we're nearly
halfway through the lecture,

00:36:42.048 --> 00:36:43.300 align:middle line:90%
so maybe we're doing OK.

00:36:43.300 --> 00:36:45.800 align:middle line:84%
So that was a
one-dimensional rectangle.

00:36:45.800 --> 00:36:49.740 align:middle line:84%
But remember that we need to be
able to approximate a rectangle

00:36:49.740 --> 00:36:51.900 align:middle line:90%
with a hyperrectangle.

00:36:51.900 --> 00:36:54.200 align:middle line:84%
In two dimensions,
it's a kind of cube.

00:36:54.200 --> 00:36:56.680 align:middle line:84%
And in three dimensions, it's
a hypercube kind of thing.

00:36:56.680 --> 00:36:57.180 align:middle line:90%
OK.

00:36:57.180 --> 00:37:00.458 align:middle line:90%
But there's another trick.

00:37:00.458 --> 00:37:02.750 align:middle line:84%
There seems like there's a
lot of tricks going on here.

00:37:02.750 --> 00:37:04.375 align:middle line:84%
But another trick
is, we basically take

00:37:04.375 --> 00:37:07.992 align:middle line:84%
a hyperrectangle that's
aligned with one axis,

00:37:07.992 --> 00:37:09.450 align:middle line:84%
and we take a
hyperrectangle that's

00:37:09.450 --> 00:37:11.070 align:middle line:90%
aligned with the other axis.

00:37:11.070 --> 00:37:14.210 align:middle line:84%
And we realize that
only in the place

00:37:14.210 --> 00:37:17.010 align:middle line:84%
where they're both kind of on,
if that makes sense, where they

00:37:17.010 --> 00:37:20.170 align:middle line:84%
cross each other, the height of
the thing, if we just add them,

00:37:20.170 --> 00:37:21.330 align:middle line:90%
is 2.

00:37:21.330 --> 00:37:24.690 align:middle line:84%
And then the height of the thing
where only one of them on is 1.

00:37:24.690 --> 00:37:26.670 align:middle line:84%
So do you see
that's this picture?

00:37:26.670 --> 00:37:30.530 align:middle line:84%
We add them and we get
this plus-shaped surface,

00:37:30.530 --> 00:37:32.810 align:middle line:84%
where in the very
middle, it's height

00:37:32.810 --> 00:37:37.210 align:middle line:84%
2 and in the other places, it's
either height 1 or height 0.

00:37:37.210 --> 00:37:39.850 align:middle line:84%
And so the observation
is that we can then

00:37:39.850 --> 00:37:43.070 align:middle line:84%
add these rectangles
in one dimension,

00:37:43.070 --> 00:37:48.010 align:middle line:84%
subtract 1, in this case,
or in general d minus 1

00:37:48.010 --> 00:37:49.550 align:middle line:90%
and bring the surface down.

00:37:49.550 --> 00:37:53.010 align:middle line:84%
And then the only place which
is kind of coming above the axis

00:37:53.010 --> 00:37:55.835 align:middle line:84%
is the place where
they're all intersecting.

00:37:55.835 --> 00:37:57.460 align:middle line:84%
And then we're just
going to threshold.

00:37:57.460 --> 00:37:59.860 align:middle line:84%
We're just going to apply
ReLU on top of that.

00:37:59.860 --> 00:38:04.040 align:middle line:84%
And it's just going to pick out
the place where they're all on.

00:38:04.040 --> 00:38:06.680 align:middle line:90%
Does that make sense?

00:38:06.680 --> 00:38:10.743 align:middle line:84%
So they'll only all be on
where they all intersect.

00:38:10.743 --> 00:38:13.160 align:middle line:84%
And then we just shift the
whole thing down and threshold,

00:38:13.160 --> 00:38:15.080 align:middle line:90%
just to slice off the top.

00:38:15.080 --> 00:38:17.780 align:middle line:84%
And it's going to give us that
little cube by just slicing off.

00:38:17.780 --> 00:38:20.760 align:middle line:84%
And we can do all of that just
by adding and doing a ReLU.

00:38:20.760 --> 00:38:22.060 align:middle line:90%
So that's the logic.

00:38:22.060 --> 00:38:24.760 align:middle line:90%


00:38:24.760 --> 00:38:29.840 align:middle line:84%
And now we're ready to
assemble all the pieces.

00:38:29.840 --> 00:38:34.680 align:middle line:84%
So the top thing is the
two-layer ReLU network that

00:38:34.680 --> 00:38:36.920 align:middle line:90%
can approximate a rectangle.

00:38:36.920 --> 00:38:42.460 align:middle line:84%
We can then sum these
things over dimensions,

00:38:42.460 --> 00:38:46.860 align:middle line:84%
subtract d minus 1 like we just
said, and then take a ReLU.

00:38:46.860 --> 00:38:50.240 align:middle line:84%
And this will now give us the
hyperrectangle that we always

00:38:50.240 --> 00:38:52.620 align:middle line:90%
wanted with this constant c.

00:38:52.620 --> 00:38:55.465 align:middle line:90%


00:38:55.465 --> 00:38:57.590 align:middle line:84%
c is the thing that we're
going to take to infinity

00:38:57.590 --> 00:39:00.730 align:middle line:90%
to make the slopes slopier.

00:39:00.730 --> 00:39:04.830 align:middle line:90%


00:39:04.830 --> 00:39:07.350 align:middle line:84%
And then finally,
we're just going

00:39:07.350 --> 00:39:12.786 align:middle line:84%
to take linear combination
of these things

00:39:12.786 --> 00:39:14.272 align:middle line:84%
with constants
alpha i, which are

00:39:14.272 --> 00:39:16.230 align:middle line:84%
going to set the height
of each hyperrectangle.

00:39:16.230 --> 00:39:20.830 align:middle line:84%
That corresponds to
adding one more layer

00:39:20.830 --> 00:39:22.090 align:middle line:90%
on top of our network.

00:39:22.090 --> 00:39:25.510 align:middle line:84%
So you can see that
in this construction,

00:39:25.510 --> 00:39:28.710 align:middle line:84%
there's a ReLU here and
there's a ReLU here.

00:39:28.710 --> 00:39:30.590 align:middle line:84%
So it really is a
three-layer network,

00:39:30.590 --> 00:39:32.750 align:middle line:90%
by the definition that we gave.

00:39:32.750 --> 00:39:35.550 align:middle line:84%
And then the final step is just
to let this constant c grow

00:39:35.550 --> 00:39:37.750 align:middle line:84%
really large so
that we really do

00:39:37.750 --> 00:39:40.350 align:middle line:90%
approximate the hyperrectangles.

00:39:40.350 --> 00:39:42.930 align:middle line:84%
And we can approximate
our arbitrary surface.

00:39:42.930 --> 00:39:45.550 align:middle line:90%


00:39:45.550 --> 00:39:48.430 align:middle line:84%
Let's just recap
where we got to.

00:39:48.430 --> 00:39:50.970 align:middle line:84%
We were going to try
to prove this theorem.

00:39:50.970 --> 00:39:53.580 align:middle line:90%


00:39:53.580 --> 00:39:55.220 align:middle line:84%
And the general
idea was we're going

00:39:55.220 --> 00:39:56.970 align:middle line:84%
to approximate these--
you could call them

00:39:56.970 --> 00:39:59.060 align:middle line:84%
bumps or these hyperrectangle
functions-- we're

00:39:59.060 --> 00:40:01.680 align:middle line:84%
going to use ReLU to
approximate those things.

00:40:01.680 --> 00:40:03.340 align:middle line:84%
And then the final
step is going to be

00:40:03.340 --> 00:40:04.980 align:middle line:90%
to take a linear combination.

00:40:04.980 --> 00:40:07.780 align:middle line:90%


00:40:07.780 --> 00:40:10.280 align:middle line:84%
So I'm going to dwell
here a little longer.

00:40:10.280 --> 00:40:13.820 align:middle line:84%
The intention of this slide
was to think about some

00:40:13.820 --> 00:40:16.620 align:middle line:84%
of the limitations
of what we just did.

00:40:16.620 --> 00:40:20.260 align:middle line:84%
I feel like some of them you
already raised like throughout.

00:40:20.260 --> 00:40:25.220 align:middle line:84%
But does anyone want to
discuss anything or point out?

00:40:25.220 --> 00:40:26.280 align:middle line:90%
Oh, yeah, go ahead.

00:40:26.280 --> 00:40:28.380 align:middle line:84%
AUDIENCE: So the two
different layers of ReLUs

00:40:28.380 --> 00:40:31.820 align:middle line:84%
do very different
functions over that?

00:40:31.820 --> 00:40:36.122 align:middle line:90%


00:40:36.122 --> 00:40:37.580 align:middle line:84%
JEREMY BERNSTEIN:
Yeah, one of them

00:40:37.580 --> 00:40:40.860 align:middle line:84%
is about approximating
the rectangle

00:40:40.860 --> 00:40:42.720 align:middle line:90%
to get the slopes of the sides.

00:40:42.720 --> 00:40:44.820 align:middle line:84%
And the other ReLU is
more about slicing off

00:40:44.820 --> 00:40:47.180 align:middle line:90%
the top of the hyperrectangle.

00:40:47.180 --> 00:40:48.060 align:middle line:90%
AUDIENCE: OK, great.

00:40:48.060 --> 00:40:50.692 align:middle line:90%


00:40:50.692 --> 00:40:51.650 align:middle line:90%
JEREMY BERNSTEIN: Yeah?

00:40:51.650 --> 00:40:53.192 align:middle line:84%
AUDIENCE: Has there
been any research

00:40:53.192 --> 00:40:56.650 align:middle line:84%
in determining if ReLU
networks actually do something?

00:40:56.650 --> 00:40:57.970 align:middle line:90%
JEREMY BERNSTEIN: Well, yeah.

00:40:57.970 --> 00:40:59.637 align:middle line:84%
Let's talk about that
on the next slide,

00:40:59.637 --> 00:41:01.290 align:middle line:90%
because I'm pretty doubtful.

00:41:01.290 --> 00:41:04.070 align:middle line:84%
It seems like quite a toy
construction-- not a toy,

00:41:04.070 --> 00:41:06.830 align:middle line:84%
but it has a very
deliberate intention,

00:41:06.830 --> 00:41:09.630 align:middle line:84%
which is proving this
approximation thing.

00:41:09.630 --> 00:41:12.392 align:middle line:84%
And a lot of the steps
seem a bit kind of fishy,

00:41:12.392 --> 00:41:13.850 align:middle line:84%
in terms of would
this be something

00:41:13.850 --> 00:41:15.350 align:middle line:84%
that a neural network
actually does,

00:41:15.350 --> 00:41:18.433 align:middle line:84%
like letting the constant
c tend to infinity?

00:41:18.433 --> 00:41:20.350 align:middle line:84%
Usually you don't want
the weights to blow up.

00:41:20.350 --> 00:41:22.142 align:middle line:84%
That would be a sign
that something's going

00:41:22.142 --> 00:41:24.410 align:middle line:90%
wrong in your neural network.

00:41:24.410 --> 00:41:27.330 align:middle line:84%
So I wanted to list
here a couple of these.

00:41:27.330 --> 00:41:32.730 align:middle line:84%
"c tends to infinity"
seems a bit questionable.

00:41:32.730 --> 00:41:33.290 align:middle line:90%
Yeah?

00:41:33.290 --> 00:41:37.010 align:middle line:84%
AUDIENCE: [INAUDIBLE] going
through the derivation

00:41:37.010 --> 00:41:38.110 align:middle line:90%
[INAUDIBLE].

00:41:38.110 --> 00:41:40.970 align:middle line:84%
So a few slides
back, you set f of x

00:41:40.970 --> 00:41:43.293 align:middle line:84%
equal to the sum of
all these rectangles?

00:41:43.293 --> 00:41:44.210 align:middle line:90%
JEREMY BERNSTEIN: Yes.

00:41:44.210 --> 00:41:46.430 align:middle line:84%
AUDIENCE: But also
from my understanding,

00:41:46.430 --> 00:41:48.450 align:middle line:90%
at this close approximate g.

00:41:48.450 --> 00:41:52.320 align:middle line:84%
And so I was confused how
the sum of the rectangles

00:41:52.320 --> 00:41:56.105 align:middle line:84%
is associated with
this [INAUDIBLE].

00:41:56.105 --> 00:41:56.980 align:middle line:90%
Does that make sense?

00:41:56.980 --> 00:41:57.897 align:middle line:90%
JEREMY BERNSTEIN: Yes.

00:41:57.897 --> 00:41:59.740 align:middle line:84%
It's because of
the use of these.

00:41:59.740 --> 00:42:02.500 align:middle line:90%


00:42:02.500 --> 00:42:04.840 align:middle line:84%
We're thinking about
it symbolically,

00:42:04.840 --> 00:42:07.560 align:middle line:84%
as being a weighted sum
of indicator functions.

00:42:07.560 --> 00:42:09.600 align:middle line:84%
So you think of the
rectangle as being

00:42:09.600 --> 00:42:13.880 align:middle line:84%
a function, which is 1 if the
input lies in this interval,

00:42:13.880 --> 00:42:15.640 align:middle line:90%
and 0 otherwise.

00:42:15.640 --> 00:42:17.960 align:middle line:90%
AUDIENCE: I see.

00:42:17.960 --> 00:42:21.820 align:middle line:84%
So I thought it said i and
not the width of an indicator.

00:42:21.820 --> 00:42:23.072 align:middle line:90%
That clarifies it.

00:42:23.072 --> 00:42:25.280 align:middle line:84%
JEREMY BERNSTEIN: Yeah,
sorry, I didn't spell it out.

00:42:25.280 --> 00:42:28.440 align:middle line:84%
There's an annotated version of
the slides on the website where

00:42:28.440 --> 00:42:30.180 align:middle line:90%
it's a bit more spelled out.

00:42:30.180 --> 00:42:33.140 align:middle line:84%
But exactly, it's not
an interval there.

00:42:33.140 --> 00:42:34.220 align:middle line:90%
It's an indicator.

00:42:34.220 --> 00:42:36.240 align:middle line:84%
AUDIENCE: And so if
it's an indicator,

00:42:36.240 --> 00:42:42.002 align:middle line:84%
then we're summing the height
at each point in the rectangle?

00:42:42.002 --> 00:42:42.960 align:middle line:90%
JEREMY BERNSTEIN: Yeah.

00:42:42.960 --> 00:42:46.760 align:middle line:84%
The way to think about
it is, for a given input,

00:42:46.760 --> 00:42:49.790 align:middle line:84%
let's say this is our
input, only this indicator

00:42:49.790 --> 00:42:51.250 align:middle line:90%
is triggered.

00:42:51.250 --> 00:42:53.010 align:middle line:84%
And then it gets
scaled by its height.

00:42:53.010 --> 00:42:55.510 align:middle line:84%
So there's a sum of terms,
but only one of them

00:42:55.510 --> 00:42:57.383 align:middle line:90%
will ever be active.

00:42:57.383 --> 00:42:58.050 align:middle line:90%
AUDIENCE: I see.

00:42:58.050 --> 00:43:00.450 align:middle line:90%
That makes sense.

00:43:00.450 --> 00:43:02.610 align:middle line:84%
JEREMY BERNSTEIN: OK, one
last question, please.

00:43:02.610 --> 00:43:05.110 align:middle line:84%
AUDIENCE: How did you
get the 4D [INAUDIBLE]?

00:43:05.110 --> 00:43:09.790 align:middle line:84%
JEREMY BERNSTEIN: Yeah, the 4D
is about counting the number

00:43:09.790 --> 00:43:13.090 align:middle line:90%
of neurons in this construction.

00:43:13.090 --> 00:43:17.410 align:middle line:84%
And the point is that in the
top purple bit, the rectangle,

00:43:17.410 --> 00:43:19.430 align:middle line:90%
there's four neurons.

00:43:19.430 --> 00:43:23.930 align:middle line:84%
Then, in the hyperrectangle, I'm
summing over d of those terms,

00:43:23.930 --> 00:43:26.590 align:middle line:90%
so it gets 4 multiplied by d.

00:43:26.590 --> 00:43:30.130 align:middle line:84%
And then so each hyper
rectangle needs 4d neurons.

00:43:30.130 --> 00:43:32.390 align:middle line:84%
And then I have the
final term, which

00:43:32.390 --> 00:43:34.370 align:middle line:84%
is the number of
hyperrectangles I need.

00:43:34.370 --> 00:43:36.270 align:middle line:90%
So that's why it's 4d.

00:43:36.270 --> 00:43:38.910 align:middle line:84%
Yeah, I hope it's right because
I just tried to work out

00:43:38.910 --> 00:43:40.570 align:middle line:90%
the constant yesterday.

00:43:40.570 --> 00:43:42.510 align:middle line:90%
But let's see.

00:43:42.510 --> 00:43:45.510 align:middle line:90%
OK, let's just go ahead.

00:43:45.510 --> 00:43:48.400 align:middle line:90%
So I wanted to think of it.

00:43:48.400 --> 00:43:50.820 align:middle line:84%
Is this really realistic
of what would actually

00:43:50.820 --> 00:43:51.717 align:middle line:90%
happen in training?

00:43:51.717 --> 00:43:53.300 align:middle line:84%
One of the questions
is, if we're just

00:43:53.300 --> 00:43:55.140 align:middle line:84%
trying to say that a
ReLU thing is going

00:43:55.140 --> 00:43:57.620 align:middle line:84%
to approximate a rectangle,
then why wouldn't we

00:43:57.620 --> 00:44:00.980 align:middle line:84%
just use a rectangle as our
basis function to begin with?

00:44:00.980 --> 00:44:05.300 align:middle line:84%
Because there's something kind
of circuitous about not just

00:44:05.300 --> 00:44:06.200 align:middle line:90%
doing that.

00:44:06.200 --> 00:44:08.260 align:middle line:84%
So let's just pretend
that we've created

00:44:08.260 --> 00:44:11.260 align:middle line:84%
a basis of these rectangles
and that someone gave us

00:44:11.260 --> 00:44:14.600 align:middle line:84%
some data, some x's and y's,
and we're going to fit it.

00:44:14.600 --> 00:44:16.940 align:middle line:84%
And just think about how
the training would go.

00:44:16.940 --> 00:44:19.220 align:middle line:84%
And what I'm claiming
is basically,

00:44:19.220 --> 00:44:21.140 align:middle line:84%
every time we see
this data point,

00:44:21.140 --> 00:44:24.980 align:middle line:84%
this rectangle would get
drawn down towards it.

00:44:24.980 --> 00:44:27.562 align:middle line:84%
And every time we see
this one, this one would.

00:44:27.562 --> 00:44:29.520 align:middle line:84%
And every time we see
this one, this one would.

00:44:29.520 --> 00:44:31.600 align:middle line:84%
And every time we see
this one, we get this one.

00:44:31.600 --> 00:44:36.060 align:middle line:84%
But there's all these other ones
which are never actually going

00:44:36.060 --> 00:44:37.340 align:middle line:90%
to move.

00:44:37.340 --> 00:44:42.180 align:middle line:84%
And so it's really
kind of unnatural.

00:44:42.180 --> 00:44:44.780 align:middle line:84%
Visually, the obvious way
to approximate this data

00:44:44.780 --> 00:44:46.550 align:middle line:84%
is just to draw a
line through it,

00:44:46.550 --> 00:44:49.125 align:middle line:84%
not to take rectangles
at each point

00:44:49.125 --> 00:44:50.750 align:middle line:84%
and move them up to
the height of the--

00:44:50.750 --> 00:44:52.790 align:middle line:84%
So there's something
unnatural about it.

00:44:52.790 --> 00:44:55.170 align:middle line:84%
But this is the
construction that we

00:44:55.170 --> 00:44:57.693 align:middle line:90%
use to prove the theorem.

00:44:57.693 --> 00:44:59.610 align:middle line:84%
In particular, the thing
I'm pointing out here

00:44:59.610 --> 00:45:03.223 align:middle line:84%
is that it seems like if you
actually fit data this way,

00:45:03.223 --> 00:45:05.389 align:middle line:84%
the training performance
is going to be really good,

00:45:05.389 --> 00:45:06.931 align:middle line:84%
but the generalization
performance is

00:45:06.931 --> 00:45:07.890 align:middle line:90%
going to be really bad.

00:45:07.890 --> 00:45:13.350 align:middle line:84%
Because you're just fitting
the data on tiny little strips,

00:45:13.350 --> 00:45:15.550 align:middle line:90%
which doesn't seem amazing.

00:45:15.550 --> 00:45:19.450 align:middle line:90%


00:45:19.450 --> 00:45:22.330 align:middle line:84%
OK, just to reassess
where we're up to,

00:45:22.330 --> 00:45:24.930 align:middle line:84%
the goal was just to
introduce you to one

00:45:24.930 --> 00:45:28.890 align:middle line:84%
particular function
approximation result,

00:45:28.890 --> 00:45:29.990 align:middle line:90%
of which there are many.

00:45:29.990 --> 00:45:31.750 align:middle line:84%
There are many
papers on this topic.

00:45:31.750 --> 00:45:33.970 align:middle line:84%
And there are some classic
papers on this topic.

00:45:33.970 --> 00:45:37.490 align:middle line:84%
So one is called Barron's
theorem, which uses something

00:45:37.490 --> 00:45:40.530 align:middle line:84%
about the Fourier
representation of the function

00:45:40.530 --> 00:45:42.050 align:middle line:90%
that you're trying to fit.

00:45:42.050 --> 00:45:46.160 align:middle line:84%
And there's a classic result.
So our construction there

00:45:46.160 --> 00:45:47.800 align:middle line:90%
was using three layers.

00:45:47.800 --> 00:45:50.193 align:middle line:84%
There's a classic result
which uses two layers

00:45:50.193 --> 00:45:52.360 align:middle line:84%
to do universal function
approximation, and actually

00:45:52.360 --> 00:45:56.360 align:middle line:84%
a more powerful result,
from my understanding.

00:45:56.360 --> 00:46:00.517 align:middle line:84%
And it uses something called the
Stone-Weierstrass theorem, which

00:46:00.517 --> 00:46:01.600 align:middle line:90%
I thought was interesting.

00:46:01.600 --> 00:46:04.120 align:middle line:84%
Because Weierstrass was
this mathematician who,

00:46:04.120 --> 00:46:07.780 align:middle line:84%
I think he was alive in the
1800s, but I'm not 100% sure.

00:46:07.780 --> 00:46:09.960 align:middle line:84%
But he came up with
that fractal curve

00:46:09.960 --> 00:46:13.020 align:middle line:84%
and was thinking about
these analysis questions.

00:46:13.020 --> 00:46:16.320 align:middle line:84%
But he was actually, seemingly
thinking about approximation

00:46:16.320 --> 00:46:16.860 align:middle line:90%
as well.

00:46:16.860 --> 00:46:21.760 align:middle line:84%
And he had a theorem, which was
like, you can use polynomials

00:46:21.760 --> 00:46:25.200 align:middle line:84%
to approximate any continuous
function or something like this.

00:46:25.200 --> 00:46:26.600 align:middle line:84%
And it's just kind
of interesting

00:46:26.600 --> 00:46:29.520 align:middle line:84%
that this branch of math
is related to a very

00:46:29.520 --> 00:46:32.640 align:middle line:90%
classical branch of math.

00:46:32.640 --> 00:46:35.480 align:middle line:90%
But that's all I have to say.

00:46:35.480 --> 00:46:38.360 align:middle line:84%
So on this slide,
I want to ask--

00:46:38.360 --> 00:46:42.270 align:middle line:84%
is the thing that we just
spent 45 minutes talking about,

00:46:42.270 --> 00:46:44.710 align:middle line:90%
is it actually important?

00:46:44.710 --> 00:46:48.230 align:middle line:84%
And here I'm just abbreviating
UFA to be universal function

00:46:48.230 --> 00:46:49.090 align:middle line:90%
approximation.

00:46:49.090 --> 00:46:50.070 align:middle line:90%
Uh-huh?

00:46:50.070 --> 00:46:51.590 align:middle line:90%
AUDIENCE: Where does this lie--

00:46:51.590 --> 00:46:54.430 align:middle line:84%
I guess the universal function
approximation theorem--

00:46:54.430 --> 00:46:58.910 align:middle line:84%
lie in accordance with using
Taylor's approximation--

00:46:58.910 --> 00:47:00.670 align:middle line:84%
or Taylor's expansion
of a function

00:47:00.670 --> 00:47:03.750 align:middle line:84%
and fitting a neural
network using polynomials

00:47:03.750 --> 00:47:04.813 align:middle line:90%
for [INAUDIBLE] function?

00:47:04.813 --> 00:47:06.230 align:middle line:84%
JEREMY BERNSTEIN:
Well, that would

00:47:06.230 --> 00:47:09.830 align:middle line:84%
be another way of just fitting
the Taylor expansion up

00:47:09.830 --> 00:47:13.190 align:middle line:84%
to some degree, which is a
bit like that Weierstrass

00:47:13.190 --> 00:47:17.050 align:middle line:84%
thing of using polynomials
to fit a continuous function.

00:47:17.050 --> 00:47:19.030 align:middle line:84%
It's like a valid
mathematical way

00:47:19.030 --> 00:47:22.310 align:middle line:84%
to approximate a function,
using a Taylor series.

00:47:22.310 --> 00:47:24.790 align:middle line:84%
But it's just not what we're
doing in deep learning.

00:47:24.790 --> 00:47:26.490 align:middle line:84%
In deep learning, we
take a neural net.

00:47:26.490 --> 00:47:29.710 align:middle line:90%
It's just different.

00:47:29.710 --> 00:47:32.370 align:middle line:84%
Thinking about that question
more, I think is great.

00:47:32.370 --> 00:47:37.030 align:middle line:84%
I don't have any
really great comments.

00:47:37.030 --> 00:47:39.052 align:middle line:84%
Yeah, putting all of
these different ways

00:47:39.052 --> 00:47:41.260 align:middle line:84%
of thinking about things
and trying to reconcile them

00:47:41.260 --> 00:47:45.420 align:middle line:84%
is like a really productive
thing to try to do.

00:47:45.420 --> 00:47:48.460 align:middle line:84%
So the point of here was to
say, is universal function

00:47:48.460 --> 00:47:49.680 align:middle line:90%
approximation sufficient?

00:47:49.680 --> 00:47:52.557 align:middle line:84%
Is it like, OK,
that was lecture 3--

00:47:52.557 --> 00:47:55.140 align:middle line:84%
now for the rest of the class,
we'll just do applications now,

00:47:55.140 --> 00:47:57.240 align:middle line:84%
because we understand the
theory of deep learning?

00:47:57.240 --> 00:47:59.540 align:middle line:90%
Probably not.

00:47:59.540 --> 00:48:01.580 align:middle line:84%
So I wanted to give
examples of other things

00:48:01.580 --> 00:48:04.080 align:middle line:84%
which are universal
function approximations.

00:48:04.080 --> 00:48:08.900 align:middle line:84%
So other examples
would be polynomials.

00:48:08.900 --> 00:48:11.300 align:middle line:84%
Can anyone think of another
example of something

00:48:11.300 --> 00:48:13.980 align:middle line:84%
which seems like it's probably
also a universal function

00:48:13.980 --> 00:48:15.000 align:middle line:90%
approximator?

00:48:15.000 --> 00:48:18.520 align:middle line:90%


00:48:18.520 --> 00:48:19.520 align:middle line:90%
AUDIENCE: Random search.

00:48:19.520 --> 00:48:20.658 align:middle line:90%
[INAUDIBLE]

00:48:20.658 --> 00:48:21.700 align:middle line:90%
JEREMY BERNSTEIN: Random?

00:48:21.700 --> 00:48:25.340 align:middle line:84%
Yeah, how would you parameterize
the space of functions?

00:48:25.340 --> 00:48:28.040 align:middle line:84%
AUDIENCE: So in the
previous example

00:48:28.040 --> 00:48:30.180 align:middle line:84%
when we had the hypercube,
we can just again

00:48:30.180 --> 00:48:32.660 align:middle line:90%
segment it into N locations.

00:48:32.660 --> 00:48:35.980 align:middle line:84%
And then when we train, we just
get stuck at whatever point

00:48:35.980 --> 00:48:37.412 align:middle line:90%
we got right.

00:48:37.412 --> 00:48:38.370 align:middle line:90%
JEREMY BERNSTEIN: Yeah.

00:48:38.370 --> 00:48:40.252 align:middle line:84%
So here we're not
really doing training.

00:48:40.252 --> 00:48:42.710 align:middle line:84%
We just want to think about
what is the space of functions.

00:48:42.710 --> 00:48:44.710 align:middle line:84%
So the space of functions
you talked about there

00:48:44.710 --> 00:48:48.230 align:middle line:84%
is linear combinations of
these hyperrectangle bumps,

00:48:48.230 --> 00:48:49.603 align:middle line:90%
which that's valid.

00:48:49.603 --> 00:48:51.770 align:middle line:84%
It's a bit wordy, so I'm
not going to write it down,

00:48:51.770 --> 00:48:53.530 align:middle line:90%
but that's valid.

00:48:53.530 --> 00:48:55.090 align:middle line:84%
The other one I
wanted to point out

00:48:55.090 --> 00:49:00.810 align:middle line:84%
is something like the
space of Python programs.

00:49:00.810 --> 00:49:02.910 align:middle line:84%
So programming languages
are, in a sense,

00:49:02.910 --> 00:49:04.750 align:middle line:84%
a universal function
approximator.

00:49:04.750 --> 00:49:07.770 align:middle line:84%
Does that mean that
we've solved machine

00:49:07.770 --> 00:49:09.690 align:middle line:90%
learning because program--

00:49:09.690 --> 00:49:12.090 align:middle line:84%
it's like many, many things
are universal function

00:49:12.090 --> 00:49:13.630 align:middle line:90%
approximators.

00:49:13.630 --> 00:49:15.022 align:middle line:90%
It's just a piece of the puzzle.

00:49:15.022 --> 00:49:16.730 align:middle line:84%
That's the thing I'm
trying to point out.

00:49:16.730 --> 00:49:18.310 align:middle line:84%
Just because you've
installed Python,

00:49:18.310 --> 00:49:20.710 align:middle line:84%
you can't start doing
machine learning.

00:49:20.710 --> 00:49:22.270 align:middle line:90%
You need something else.

00:49:22.270 --> 00:49:26.010 align:middle line:90%


00:49:26.010 --> 00:49:29.750 align:middle line:84%
And is it necessary for
machine learning to work?

00:49:29.750 --> 00:49:31.650 align:middle line:84%
So if your machine
learning model

00:49:31.650 --> 00:49:33.710 align:middle line:84%
is not a universal
function approximator,

00:49:33.710 --> 00:49:34.830 align:middle line:90%
should you be worried?

00:49:34.830 --> 00:49:38.960 align:middle line:90%


00:49:38.960 --> 00:49:40.120 align:middle line:90%
Someone's saying no.

00:49:40.120 --> 00:49:44.400 align:middle line:84%
Yeah, I sort of agree,
but I'm honestly not sure.

00:49:44.400 --> 00:49:47.220 align:middle line:84%
Yeah, I think it's just an
interesting question to ask.

00:49:47.220 --> 00:49:48.380 align:middle line:90%
Do we need it to be?

00:49:48.380 --> 00:49:49.583 align:middle line:90%
My guess is not.

00:49:49.583 --> 00:49:51.000 align:middle line:84%
My guess is probably
we don't need

00:49:51.000 --> 00:49:52.560 align:middle line:84%
to have a universal
function approximator

00:49:52.560 --> 00:49:53.540 align:middle line:90%
to do machine learning.

00:49:53.540 --> 00:49:59.080 align:middle line:84%
But I think it's a little
unclear how to answer that.

00:49:59.080 --> 00:50:06.040 align:middle line:84%
So now I want to move on to
the second half of the lecture,

00:50:06.040 --> 00:50:09.080 align:middle line:84%
which is going to
be thinking more

00:50:09.080 --> 00:50:13.160 align:middle line:84%
about that question
of width versus depth.

00:50:13.160 --> 00:50:15.880 align:middle line:84%
So again, would you rather
scale the width of your model

00:50:15.880 --> 00:50:17.260 align:middle line:90%
or scale the depth?

00:50:17.260 --> 00:50:20.000 align:middle line:90%


00:50:20.000 --> 00:50:23.160 align:middle line:84%
And again, we've just tried
to prove in some sense

00:50:23.160 --> 00:50:26.840 align:middle line:84%
that with a very small
depth network, like depth 3,

00:50:26.840 --> 00:50:28.740 align:middle line:90%
we can fit any function.

00:50:28.740 --> 00:50:31.700 align:middle line:84%
So that's coming back
to this question of,

00:50:31.700 --> 00:50:35.040 align:middle line:84%
if three layers is enough,
why would we ever want

00:50:35.040 --> 00:50:37.490 align:middle line:90%
to have a deep network?

00:50:37.490 --> 00:50:41.430 align:middle line:90%
Does anyone have a thought?

00:50:41.430 --> 00:50:42.250 align:middle line:90%
What do you think?

00:50:42.250 --> 00:50:43.430 align:middle line:90%
Yeah?

00:50:43.430 --> 00:50:47.270 align:middle line:84%
AUDIENCE: I'm guessing because
in the terms of nailing

00:50:47.270 --> 00:50:52.310 align:middle line:84%
the network and stuff, you
can approximate, I'm guessing,

00:50:52.310 --> 00:50:56.350 align:middle line:84%
some sort of arbitrary
nonlinearity, because you have

00:50:56.350 --> 00:50:58.570 align:middle line:84%
pointwise nonlinearities
in each of the neurons,

00:50:58.570 --> 00:50:59.930 align:middle line:90%
you have function compositions.

00:50:59.930 --> 00:51:05.190 align:middle line:84%
So those would help you
to do a complex mapping.

00:51:05.190 --> 00:51:09.150 align:middle line:84%
But width would be going back
to the universal approximation

00:51:09.150 --> 00:51:10.150 align:middle line:90%
theorem.

00:51:10.150 --> 00:51:11.150 align:middle line:90%
JEREMY BERNSTEIN: Right.

00:51:11.150 --> 00:51:11.710 align:middle line:90%
That's great.

00:51:11.710 --> 00:51:13.250 align:middle line:90%
Let me just summarize.

00:51:13.250 --> 00:51:13.750 align:middle line:90%
I forgot.

00:51:13.750 --> 00:51:16.410 align:middle line:84%
I'm always supposed to summarize
questions and comments,

00:51:16.410 --> 00:51:17.410 align:middle line:90%
but I forgot about that.

00:51:17.410 --> 00:51:21.910 align:middle line:84%
But basically, if we
layer lots of layers,

00:51:21.910 --> 00:51:24.250 align:middle line:84%
then we have a kind
of compound effect.

00:51:24.250 --> 00:51:26.010 align:middle line:90%
We have more nonlinearities.

00:51:26.010 --> 00:51:27.510 align:middle line:84%
It seems like maybe
that can give us

00:51:27.510 --> 00:51:31.790 align:middle line:84%
a richer space of functions
through compositionality.

00:51:31.790 --> 00:51:33.010 align:middle line:90%
Yeah, that's a great point.

00:51:33.010 --> 00:51:34.520 align:middle line:90%
Sorry, what was your--

00:51:34.520 --> 00:51:36.820 align:middle line:84%
AUDIENCE: Maybe on surface
what I think we can have is

00:51:36.820 --> 00:51:38.660 align:middle line:90%
we can put on lines and dots.

00:51:38.660 --> 00:51:43.860 align:middle line:84%
The width might be faster
or take less memory.

00:51:43.860 --> 00:51:47.155 align:middle line:84%
But if you don't, then
that might be good.

00:51:47.155 --> 00:51:48.780 align:middle line:84%
JEREMY BERNSTEIN:
This is a great point

00:51:48.780 --> 00:51:50.155 align:middle line:84%
that, if you think
about building

00:51:50.155 --> 00:51:53.760 align:middle line:84%
a machine learning system,
width can be parallelized.

00:51:53.760 --> 00:51:56.660 align:middle line:84%
It makes better use
of parallel hardware.

00:51:56.660 --> 00:51:59.940 align:middle line:84%
So from a systems perspective,
if you can get away

00:51:59.940 --> 00:52:02.100 align:middle line:84%
with a shallow network
that's very wide,

00:52:02.100 --> 00:52:06.620 align:middle line:84%
you would really love that,
because you can parallelize it.

00:52:06.620 --> 00:52:10.080 align:middle line:84%
And that's illustrative that
there's different constraints.

00:52:10.080 --> 00:52:13.900 align:middle line:84%
There's the, can I represent
complex and interesting

00:52:13.900 --> 00:52:14.458 align:middle line:90%
functions?

00:52:14.458 --> 00:52:16.500 align:middle line:84%
And then there's the
computational constraint of,

00:52:16.500 --> 00:52:18.200 align:middle line:90%
how efficient is this to run?

00:52:18.200 --> 00:52:21.280 align:middle line:84%
There's other
constraints of that sort.

00:52:21.280 --> 00:52:24.380 align:middle line:90%


00:52:24.380 --> 00:52:27.300 align:middle line:84%
So on this slide,
I wanted to just

00:52:27.300 --> 00:52:30.820 align:middle line:84%
think, what are the
advantages of width?

00:52:30.820 --> 00:52:32.830 align:middle line:90%
So one is parallelism.

00:52:32.830 --> 00:52:35.890 align:middle line:90%


00:52:35.890 --> 00:52:40.090 align:middle line:84%
One is a three-layer neural
net is a universal function

00:52:40.090 --> 00:52:43.230 align:middle line:84%
approximator, or maybe
even two layers is enough.

00:52:43.230 --> 00:52:47.970 align:middle line:90%


00:52:47.970 --> 00:52:50.670 align:middle line:84%
Can anyone think of any
other advantages of width?

00:52:50.670 --> 00:52:56.335 align:middle line:90%


00:52:56.335 --> 00:52:58.210 align:middle line:84%
AUDIENCE: You get the
feature representations

00:52:58.210 --> 00:53:00.970 align:middle line:84%
that are nicer and more
easily interpretable.

00:53:00.970 --> 00:53:05.390 align:middle line:84%
JEREMY BERNSTEIN: Mm-hmm,
easier to interpret.

00:53:05.390 --> 00:53:06.890 align:middle line:84%
I'm trying to
remember if there were

00:53:06.890 --> 00:53:08.270 align:middle line:90%
any others that I thought of.

00:53:08.270 --> 00:53:09.370 align:middle line:90%
Does anyone else have any?

00:53:09.370 --> 00:53:13.050 align:middle line:90%


00:53:13.050 --> 00:53:17.568 align:middle line:84%
AUDIENCE: It appears
[INAUDIBLE] gradients.

00:53:17.568 --> 00:53:20.110 align:middle line:84%
JEREMY BERNSTEIN: Yeah, easier
to train, that's a good point.

00:53:20.110 --> 00:53:23.730 align:middle line:90%


00:53:23.730 --> 00:53:25.110 align:middle line:90%
Easier to optimize.

00:53:25.110 --> 00:53:29.590 align:middle line:90%


00:53:29.590 --> 00:53:33.040 align:middle line:84%
Yeah, because whenever
you do compound anything,

00:53:33.040 --> 00:53:35.540 align:middle line:84%
it is very liable to
get out of control.

00:53:35.540 --> 00:53:38.160 align:middle line:84%
So making your network
too deep without being

00:53:38.160 --> 00:53:40.480 align:middle line:84%
careful about how you
do that can actually

00:53:40.480 --> 00:53:42.040 align:middle line:90%
just break the training.

00:53:42.040 --> 00:53:44.760 align:middle line:84%
So if you just take a
vanilla multilayer perceptron

00:53:44.760 --> 00:53:47.860 align:middle line:84%
and make it 50 layers deep
and just try to train it,

00:53:47.860 --> 00:53:53.320 align:middle line:84%
you'll see it's very
difficult to get it to train.

00:53:53.320 --> 00:53:54.200 align:middle line:90%
OK.

00:53:54.200 --> 00:53:57.840 align:middle line:84%
I can't think of any others,
unless anyone else had.

00:53:57.840 --> 00:54:02.200 align:middle line:84%
So the point of this slide was
to say, well, OK, obviously wide

00:54:02.200 --> 00:54:04.920 align:middle line:90%
is better than deep, right?

00:54:04.920 --> 00:54:06.060 align:middle line:90%
That's obvious now.

00:54:06.060 --> 00:54:09.660 align:middle line:84%
But then the next slide
is supposed to say, is it?

00:54:09.660 --> 00:54:14.120 align:middle line:90%


00:54:14.120 --> 00:54:17.395 align:middle line:84%
Because now in the last
part of the lecture

00:54:17.395 --> 00:54:19.020 align:middle line:84%
or almost the last
part of the lecture,

00:54:19.020 --> 00:54:20.812 align:middle line:84%
we're going to think
about something called

00:54:20.812 --> 00:54:25.320 align:middle line:84%
depth separation results, which
is another potentially quite

00:54:25.320 --> 00:54:30.750 align:middle line:84%
stylized theoretical result,
but it is quite interesting.

00:54:30.750 --> 00:54:34.750 align:middle line:84%
So remember that the universal
function approximation results

00:54:34.750 --> 00:54:37.310 align:middle line:84%
suggested that we need to
make our network exponentially

00:54:37.310 --> 00:54:40.150 align:middle line:84%
wide in order to be able
to approximate everything

00:54:40.150 --> 00:54:41.250 align:middle line:90%
that we want to.

00:54:41.250 --> 00:54:43.310 align:middle line:90%
And that seems bad.

00:54:43.310 --> 00:54:45.150 align:middle line:84%
So the point of a
depth separation result

00:54:45.150 --> 00:54:48.070 align:middle line:84%
is to say that there's
actually a deep network that

00:54:48.070 --> 00:54:52.870 align:middle line:84%
can represent a particular
function quite efficiently.

00:54:52.870 --> 00:54:54.870 align:middle line:84%
And if you tried to
represent the same function

00:54:54.870 --> 00:54:58.150 align:middle line:84%
with a three-layer network
or a shallower network,

00:54:58.150 --> 00:55:02.470 align:middle line:84%
you would need exponentially
more neurons to do that.

00:55:02.470 --> 00:55:07.910 align:middle line:84%
So the shape of a depth
separation result,

00:55:07.910 --> 00:55:12.610 align:middle line:84%
I'm just trying to break it up
into stages, is first of all,

00:55:12.610 --> 00:55:14.690 align:middle line:84%
we pick a property
of a function.

00:55:14.690 --> 00:55:17.430 align:middle line:90%


00:55:17.430 --> 00:55:19.750 align:middle line:84%
And the example we're
going to think of

00:55:19.750 --> 00:55:21.910 align:middle line:90%
is the number of linear regions.

00:55:21.910 --> 00:55:24.870 align:middle line:90%


00:55:24.870 --> 00:55:27.220 align:middle line:84%
We think of the function
as being piecewise linear.

00:55:27.220 --> 00:55:30.140 align:middle line:84%
And how many linear
pieces are there?

00:55:30.140 --> 00:55:32.460 align:middle line:84%
And then we construct
a deep network

00:55:32.460 --> 00:55:36.980 align:middle line:84%
that has a lot of that property,
but is not a particularly

00:55:36.980 --> 00:55:38.360 align:middle line:90%
big network in some sense.

00:55:38.360 --> 00:55:39.900 align:middle line:84%
It doesn't have
a lot of neurons.

00:55:39.900 --> 00:55:42.520 align:middle line:84%
And that is a
constructive thing.

00:55:42.520 --> 00:55:44.140 align:middle line:84%
We say, here is a
neural network which

00:55:44.140 --> 00:55:45.980 align:middle line:90%
has a lot of this property.

00:55:45.980 --> 00:55:48.780 align:middle line:84%
And then we say, now we're going
to prove that it's actually

00:55:48.780 --> 00:55:51.900 align:middle line:84%
impossible for a shallow network
to have that same property

00:55:51.900 --> 00:55:55.940 align:middle line:84%
unless it has exponentially
many neurons or a huge number

00:55:55.940 --> 00:55:57.700 align:middle line:90%
of neurons.

00:55:57.700 --> 00:55:59.440 align:middle line:84%
So that's the shape
of this result.

00:55:59.440 --> 00:56:03.020 align:middle line:84%
And again, there's a
literature which just proves

00:56:03.020 --> 00:56:05.340 align:middle line:84%
different varieties of
this kind of result.

00:56:05.340 --> 00:56:07.560 align:middle line:84%
And that's called depth
separation results.

00:56:07.560 --> 00:56:10.640 align:middle line:84%
We're just going to show
you one particular example.

00:56:10.640 --> 00:56:13.580 align:middle line:90%


00:56:13.580 --> 00:56:16.980 align:middle line:84%
So recall that a
piecewise linear function

00:56:16.980 --> 00:56:20.080 align:middle line:84%
is a function where if you
break it up into segments,

00:56:20.080 --> 00:56:21.600 align:middle line:90%
each segment is linear.

00:56:21.600 --> 00:56:23.200 align:middle line:84%
And then they kind
of stitch together.

00:56:23.200 --> 00:56:26.610 align:middle line:90%


00:56:26.610 --> 00:56:29.690 align:middle line:84%
And we'll think about
defining the number of kinks

00:56:29.690 --> 00:56:31.730 align:middle line:90%
in such a function to mean--

00:56:31.730 --> 00:56:34.850 align:middle line:90%


00:56:34.850 --> 00:56:38.810 align:middle line:84%
a kink is how many of
these discontinuities

00:56:38.810 --> 00:56:42.090 align:middle line:84%
are there in the derivative,
how many places does

00:56:42.090 --> 00:56:44.610 align:middle line:90%
the derivative suddenly change.

00:56:44.610 --> 00:56:49.290 align:middle line:84%
And so for this function, you
can see there's four kinks.

00:56:49.290 --> 00:56:53.210 align:middle line:84%
And that's the property
that we're going to pick.

00:56:53.210 --> 00:56:55.130 align:middle line:84%
And we're going to
show that there's

00:56:55.130 --> 00:56:57.290 align:middle line:84%
a deep network with
not very many neurons

00:56:57.290 --> 00:56:58.453 align:middle line:90%
that has a lot of kinks.

00:56:58.453 --> 00:57:00.870 align:middle line:84%
And if you want to have that
many kinks in a wide network,

00:57:00.870 --> 00:57:02.390 align:middle line:84%
you need exponentially
many neurons.

00:57:02.390 --> 00:57:05.782 align:middle line:90%


00:57:05.782 --> 00:57:07.990 align:middle line:84%
That's the strategy of what
we're going to try to do.

00:57:07.990 --> 00:57:11.310 align:middle line:84%
But the first claim is that ReLU
networks are piecewise linear.

00:57:11.310 --> 00:57:15.090 align:middle line:84%
So if you have a ReLU network
with one input and one output,

00:57:15.090 --> 00:57:17.650 align:middle line:84%
it will look something
like this green curve.

00:57:17.650 --> 00:57:22.650 align:middle line:84%
Can anyone tell me
why that's the case?

00:57:22.650 --> 00:57:24.500 align:middle line:84%
Or does anyone think
it is the case?

00:57:24.500 --> 00:57:26.420 align:middle line:84%
Or does anyone think
it's not the case?

00:57:26.420 --> 00:57:36.120 align:middle line:90%


00:57:36.120 --> 00:57:36.920 align:middle line:90%
Yeah?

00:57:36.920 --> 00:57:39.560 align:middle line:84%
AUDIENCE: It's because
the ReLU is essentially

00:57:39.560 --> 00:57:41.460 align:middle line:84%
a linear function,
but with a threshold.

00:57:41.460 --> 00:57:45.040 align:middle line:84%
So when you add them, it
creates some kind of shape.

00:57:45.040 --> 00:57:47.400 align:middle line:84%
And the transformation
that you did before

00:57:47.400 --> 00:57:51.700 align:middle line:84%
is linear, so it's only going to
distort it but in a linear way.

00:57:51.700 --> 00:57:53.400 align:middle line:90%
JEREMY BERNSTEIN: Great.

00:57:53.400 --> 00:57:55.760 align:middle line:84%
Yeah, I think with
this kind of statement,

00:57:55.760 --> 00:57:57.943 align:middle line:84%
with many sorts of
statements-- yeah, exactly.

00:57:57.943 --> 00:57:59.360 align:middle line:84%
But with many sorts
of statements,

00:57:59.360 --> 00:58:01.040 align:middle line:90%
we want to break them up into--

00:58:01.040 --> 00:58:04.200 align:middle line:90%


00:58:04.200 --> 00:58:05.700 align:middle line:84%
I'm not sure what
the right word is.

00:58:05.700 --> 00:58:08.600 align:middle line:84%
But we want to be able
to say, if f and g are

00:58:08.600 --> 00:58:12.640 align:middle line:84%
both piecewise linear, then that
implies that f composed with g

00:58:12.640 --> 00:58:14.200 align:middle line:90%
is piecewise linear.

00:58:14.200 --> 00:58:16.840 align:middle line:84%
If f and g are both
piecewise linear,

00:58:16.840 --> 00:58:20.120 align:middle line:84%
that implies that f plus
g is piecewise linear.

00:58:20.120 --> 00:58:23.150 align:middle line:84%
If f is piecewise
linear, then alpha

00:58:23.150 --> 00:58:25.050 align:middle line:90%
f is also piecewise linear.

00:58:25.050 --> 00:58:27.670 align:middle line:84%
So you break up the
construction of the function

00:58:27.670 --> 00:58:30.710 align:middle line:84%
into a series of
combinations, and you

00:58:30.710 --> 00:58:32.310 align:middle line:84%
prove that the
property is preserved

00:58:32.310 --> 00:58:35.090 align:middle line:90%
under those combinations.

00:58:35.090 --> 00:58:36.590 align:middle line:84%
And then that and
then basically you

00:58:36.590 --> 00:58:42.750 align:middle line:84%
realize that, OK, a
ReLU neural network,

00:58:42.750 --> 00:58:44.568 align:middle line:84%
the ReLU function
is piecewise linear.

00:58:44.568 --> 00:58:46.610 align:middle line:84%
And then it's just adding
and multiplying things.

00:58:46.610 --> 00:58:47.870 align:middle line:84%
So you're always
going to get something

00:58:47.870 --> 00:58:49.090 align:middle line:90%
that's piecewise linear.

00:58:49.090 --> 00:58:51.610 align:middle line:90%


00:58:51.610 --> 00:58:54.310 align:middle line:84%
So based on that, let's now
construct a depth separation

00:58:54.310 --> 00:58:56.750 align:middle line:84%
based on counting
the number of kinks.

00:58:56.750 --> 00:59:00.670 align:middle line:84%
So I'll first give
you the intuition

00:59:00.670 --> 00:59:04.270 align:middle line:84%
and then write something
a bit more formal.

00:59:04.270 --> 00:59:06.590 align:middle line:84%
But first of all, think
about the point of view

00:59:06.590 --> 00:59:08.470 align:middle line:90%
of an individual neuron.

00:59:08.470 --> 00:59:14.030 align:middle line:84%
And we're going to add
one extra input neuron.

00:59:14.030 --> 00:59:17.510 align:middle line:84%
And symbolically, then we
think about feeding a function,

00:59:17.510 --> 00:59:21.320 align:middle line:84%
like a one-dimensional
function, as inputs.

00:59:21.320 --> 00:59:23.780 align:middle line:84%
So we're going to think of
the output of this neuron, y

00:59:23.780 --> 00:59:27.940 align:middle line:84%
of x, as just being the
summation of scalar multiples

00:59:27.940 --> 00:59:31.820 align:middle line:90%
of the input functions.

00:59:31.820 --> 00:59:37.340 align:middle line:84%
The observation of this slide
is that if we add functions,

00:59:37.340 --> 00:59:41.060 align:middle line:84%
at most, we add the
number of kinks.

00:59:41.060 --> 00:59:44.020 align:middle line:84%
And I'm just going to try
to persuade you that that's

00:59:44.020 --> 00:59:45.860 align:middle line:90%
the case with a picture.

00:59:45.860 --> 00:59:47.860 align:middle line:84%
So here we're thinking
of, the green function

00:59:47.860 --> 00:59:49.580 align:middle line:90%
is our first function.

00:59:49.580 --> 00:59:52.260 align:middle line:84%
And the blue function
is our second function.

00:59:52.260 --> 00:59:54.220 align:middle line:84%
And then the purple
function shows the result

00:59:54.220 --> 00:59:55.820 align:middle line:90%
of adding them together.

00:59:55.820 --> 01:00:00.540 align:middle line:84%
And notice that the green
function has one kink.

01:00:00.540 --> 01:00:03.500 align:middle line:90%
The blue function has two kinks.

01:00:03.500 --> 01:00:05.620 align:middle line:84%
And then the purple
function, which is their sum,

01:00:05.620 --> 01:00:07.140 align:middle line:90%
has three kinks.

01:00:07.140 --> 01:00:09.180 align:middle line:84%
And nothing more, you
can't have more than that,

01:00:09.180 --> 01:00:12.680 align:middle line:84%
because you only can
have kinks at the places

01:00:12.680 --> 01:00:14.680 align:middle line:84%
where the functions that
you're adding had them.

01:00:14.680 --> 01:00:17.320 align:middle line:84%
So this is how we think
about adding functions.

01:00:17.320 --> 01:00:19.490 align:middle line:84%
We can at most add
the number of kicks.

01:00:19.490 --> 01:00:21.650 align:middle line:84%
But potentially the
kicks could happen

01:00:21.650 --> 01:00:24.355 align:middle line:84%
in the same place,
in which case, it

01:00:24.355 --> 01:00:26.730 align:middle line:84%
would be less than adding,
because that wouldn't give you

01:00:26.730 --> 01:00:27.750 align:middle line:90%
an extra one.

01:00:27.750 --> 01:00:29.890 align:middle line:84%
But generally, if they're
generally spaced out,

01:00:29.890 --> 01:00:31.970 align:middle line:90%
then they would add.

01:00:31.970 --> 01:00:36.570 align:middle line:84%
Next intuition, the
effect of applying ReLU.

01:00:36.570 --> 01:00:39.810 align:middle line:84%
So we think now that
we have a function f.

01:00:39.810 --> 01:00:42.010 align:middle line:90%
And then we just take ReLU of f.

01:00:42.010 --> 01:00:46.670 align:middle line:84%
So we think that the output
is ReLU of f of x, basically.

01:00:46.670 --> 01:00:48.810 align:middle line:90%
So y of x is ReLU of f of x.

01:00:48.810 --> 01:00:53.610 align:middle line:84%
And now the claim is that
if we apply ReLU, at most,

01:00:53.610 --> 01:00:56.950 align:middle line:90%
we double the number of kinks.

01:00:56.950 --> 01:00:59.170 align:middle line:84%
And so this picture is
supposed to illustrate.

01:00:59.170 --> 01:01:02.530 align:middle line:84%
We think about the red
function as being f.

01:01:02.530 --> 01:01:08.970 align:middle line:84%
And then if we apply ReLU, we
just chop off the positive part.

01:01:08.970 --> 01:01:10.530 align:middle line:84%
And the observation
is that if we

01:01:10.530 --> 01:01:14.810 align:middle line:84%
had a linear region, that
can at most get split up

01:01:14.810 --> 01:01:17.900 align:middle line:90%
into two linear regions.

01:01:17.900 --> 01:01:21.200 align:middle line:84%
And if you just count
on this picture,

01:01:21.200 --> 01:01:24.540 align:middle line:84%
I think you see that the
red function had 1, 2, 3, 4,

01:01:24.540 --> 01:01:32.948 align:middle line:84%
5 kinks, whereas the green
function has 1, 2, 3, 4, 5, 6,

01:01:32.948 --> 01:01:35.500 align:middle line:90%
7, 8, 9, 9 kinks.

01:01:35.500 --> 01:01:38.020 align:middle line:84%
So you can see it did
nearly double in this case.

01:01:38.020 --> 01:01:43.920 align:middle line:90%


01:01:43.920 --> 01:01:47.600 align:middle line:84%
So now to be a bit
more formal, we're

01:01:47.600 --> 01:01:50.840 align:middle line:84%
going to imagine that
we have a deep ReLU

01:01:50.840 --> 01:01:52.520 align:middle line:90%
network with many layers.

01:01:52.520 --> 01:01:57.880 align:middle line:84%
And we're going to examine
one layer inside the network.

01:01:57.880 --> 01:02:02.920 align:middle line:84%
And we're going to give a
recursive definition of what

01:02:02.920 --> 01:02:04.300 align:middle line:90%
this layer looks like.

01:02:04.300 --> 01:02:07.760 align:middle line:84%
So the function
at layer L is ReLU

01:02:07.760 --> 01:02:11.160 align:middle line:84%
of the weight matrix times the
function from the previous layer

01:02:11.160 --> 01:02:13.880 align:middle line:90%
plus the bias.

01:02:13.880 --> 01:02:18.500 align:middle line:84%
And this is going to be
an n-dimensional vector.

01:02:18.500 --> 01:02:22.800 align:middle line:84%
The weight matrix is going
to be an n by n matrix.

01:02:22.800 --> 01:02:25.780 align:middle line:84%
This is going to be an
n-dimensional vector.

01:02:25.780 --> 01:02:27.840 align:middle line:84%
And the bias is also an
n-dimensional vector.

01:02:27.840 --> 01:02:31.740 align:middle line:84%
So that's the shape of
all of these things.

01:02:31.740 --> 01:02:36.500 align:middle line:84%
And then we define
this capital let

01:02:36.500 --> 01:02:41.020 align:middle line:84%
KINKS-L to denote the maximum
number of kinks over the n

01:02:41.020 --> 01:02:43.120 align:middle line:90%
coordinates of the output.

01:02:43.120 --> 01:02:45.380 align:middle line:84%
So remember, this is an
n-dimensional vector.

01:02:45.380 --> 01:02:48.220 align:middle line:84%
Each coordinate could have
a certain number of kinks

01:02:48.220 --> 01:02:50.140 align:middle line:84%
in its function,
because each coordinate

01:02:50.140 --> 01:02:53.860 align:middle line:90%
is a one-dimensional function.

01:02:53.860 --> 01:02:57.140 align:middle line:84%
And then we define this
variable to be the maximum

01:02:57.140 --> 01:02:58.860 align:middle line:90%
over those coordinates.

01:02:58.860 --> 01:03:02.420 align:middle line:84%
And maybe I'm just going
to write it as a claim.

01:03:02.420 --> 01:03:08.620 align:middle line:84%
But the claim is that KINKS-L
can be no greater than 2 times

01:03:08.620 --> 01:03:15.090 align:middle line:84%
the width times the maximum
number of kinks in the input.

01:03:15.090 --> 01:03:17.670 align:middle line:84%
So maybe you could
think more about this.

01:03:17.670 --> 01:03:21.150 align:middle line:84%
So the 2 comes from the
doubling effect of the ReLU.

01:03:21.150 --> 01:03:25.010 align:middle line:84%
And the n comes from the fact
that you're adding n things.

01:03:25.010 --> 01:03:27.010 align:middle line:84%
And then the trick is
just to think about taking

01:03:27.010 --> 01:03:29.390 align:middle line:90%
the max over each coordinate.

01:03:29.390 --> 01:03:32.970 align:middle line:90%


01:03:32.970 --> 01:03:35.050 align:middle line:84%
And then the other
observation is

01:03:35.050 --> 01:03:38.130 align:middle line:84%
that the kinks at
the input, the input

01:03:38.130 --> 01:03:42.770 align:middle line:84%
is basically 1, because
there's no kinks.

01:03:42.770 --> 01:03:45.790 align:middle line:90%
Does that make sense?

01:03:45.790 --> 01:03:47.050 align:middle line:90%
AUDIENCE: It should be 0.

01:03:47.050 --> 01:03:48.230 align:middle line:84%
JEREMY BERNSTEIN: It
should be 0, you're right.

01:03:48.230 --> 01:03:49.250 align:middle line:90%
But it's not going to--

01:03:49.250 --> 01:03:54.962 align:middle line:90%


01:03:54.962 --> 01:03:56.670 align:middle line:84%
Let's think about this
after the lecture.

01:03:56.670 --> 01:04:02.170 align:middle line:84%
But the point is that each time
we can at most multiply by 2n.

01:04:02.170 --> 01:04:08.290 align:middle line:84%
So I'm trying to get the answer
that KINKS-L is no greater than

01:04:08.290 --> 01:04:17.740 align:middle line:84%
2n to the power L. OK,
well, if KINKS-0 is 0,

01:04:17.740 --> 01:04:19.800 align:middle line:84%
it's still less
than or equal to 1.

01:04:19.800 --> 01:04:21.440 align:middle line:90%
[LAUGHTER]

01:04:21.440 --> 01:04:24.620 align:middle line:84%
Yeah, it looks like
some accounting problem.

01:04:24.620 --> 01:04:26.380 align:middle line:84%
But you see what
the argument is.

01:04:26.380 --> 01:04:27.580 align:middle line:90%
I just need to figure out--

01:04:27.580 --> 01:04:28.080 align:middle line:90%
Yeah?

01:04:28.080 --> 01:04:29.788 align:middle line:84%
AUDIENCE: You could
start with one layer.

01:04:29.788 --> 01:04:31.372 align:middle line:84%
JEREMY BERNSTEIN:
Yeah, you could just

01:04:31.372 --> 01:04:33.180 align:middle line:84%
start with the base
case being one layer.

01:04:33.180 --> 01:04:34.560 align:middle line:90%
That's a great--

01:04:34.560 --> 01:04:36.800 align:middle line:84%
AUDIENCE: You should make
it segmented instead.

01:04:36.800 --> 01:04:40.640 align:middle line:84%
Otherwise, the ReLU adds a
kink where there was none.

01:04:40.640 --> 01:04:45.160 align:middle line:84%
Well, 0 to 1 is
more [INAUDIBLE].

01:04:45.160 --> 01:04:46.720 align:middle line:90%
JEREMY BERNSTEIN: Yes.

01:04:46.720 --> 01:04:47.560 align:middle line:90%
OK.

01:04:47.560 --> 01:04:49.500 align:middle line:90%
Let's just talk more about this.

01:04:49.500 --> 01:04:51.440 align:middle line:84%
But I think you
get the intuition

01:04:51.440 --> 01:04:52.780 align:middle line:90%
in making this rigorous.

01:04:52.780 --> 01:04:54.322 align:middle line:84%
You just need to
think about it a bit

01:04:54.322 --> 01:04:58.120 align:middle line:84%
to work out what the
base case should be.

01:04:58.120 --> 01:05:01.440 align:middle line:84%
But I hope that the
message gets across.

01:05:01.440 --> 01:05:05.500 align:middle line:84%
So in principle, we
showed this statement.

01:05:05.500 --> 01:05:09.560 align:middle line:90%


01:05:09.560 --> 01:05:13.310 align:middle line:84%
As a reminder, n is the
width, capital L is the depth,

01:05:13.310 --> 01:05:18.890 align:middle line:84%
and KINKS-L should be the
maximum over the coordinates.

01:05:18.890 --> 01:05:22.910 align:middle line:90%


01:05:22.910 --> 01:05:25.790 align:middle line:84%
And so the thing we
wanted to point out

01:05:25.790 --> 01:05:34.530 align:middle line:84%
is that the upper bound grows
at most polynomially in width,

01:05:34.530 --> 01:05:37.870 align:middle line:90%
but exponentially in depth.

01:05:37.870 --> 01:05:41.510 align:middle line:84%
So this is suggestive that
there's a big benefit,

01:05:41.510 --> 01:05:43.710 align:middle line:84%
if we're trying to maximize
the number of kinks,

01:05:43.710 --> 01:05:45.230 align:middle line:90%
in making the network deeper.

01:05:45.230 --> 01:05:46.650 align:middle line:84%
Because every time
we add a layer,

01:05:46.650 --> 01:05:50.330 align:middle line:84%
we double things, rather
than making it wider.

01:05:50.330 --> 01:05:53.550 align:middle line:84%
So if our goal is to be able
to approximate functions

01:05:53.550 --> 01:05:56.990 align:middle line:84%
with many kinks,
then it's better

01:05:56.990 --> 01:06:01.750 align:middle line:84%
to make the network
deeper rather than wider.

01:06:01.750 --> 01:06:05.530 align:middle line:84%
Does anyone have an objection
to what we've shown?

01:06:05.530 --> 01:06:09.360 align:middle line:84%
Because it is just
an upper bound.

01:06:09.360 --> 01:06:10.860 align:middle line:84%
So the objection
that you could have

01:06:10.860 --> 01:06:12.880 align:middle line:84%
is like, that's
just an upper bound.

01:06:12.880 --> 01:06:14.740 align:middle line:90%
Maybe it's never attained.

01:06:14.740 --> 01:06:17.242 align:middle line:84%
Maybe in practice, when you
actually build a very deep ReLU

01:06:17.242 --> 01:06:18.700 align:middle line:84%
network, it's very
difficult to get

01:06:18.700 --> 01:06:20.900 align:middle line:84%
there to be a lot of kinks
in the output function.

01:06:20.900 --> 01:06:26.220 align:middle line:84%
So the next step is to
provide a construction

01:06:26.220 --> 01:06:28.680 align:middle line:84%
to show you that, no,
this can actually happen.

01:06:28.680 --> 01:06:31.280 align:middle line:84%
And again, I'm just
going to skip the slide.

01:06:31.280 --> 01:06:36.060 align:middle line:84%
So again, the argument
is that this g defined

01:06:36.060 --> 01:06:39.080 align:middle line:84%
as a really small ReLU
network in this form,

01:06:39.080 --> 01:06:41.100 align:middle line:84%
if you plot what that
actually amounts to,

01:06:41.100 --> 01:06:43.980 align:middle line:90%
it amounts to a triangle.

01:06:43.980 --> 01:06:47.440 align:middle line:84%
And the property of a triangle
is that if you iterate it,

01:06:47.440 --> 01:06:50.900 align:middle line:84%
you apply it to itself-- so
you do g composed with g,

01:06:50.900 --> 01:06:52.700 align:middle line:90%
you actually get two triangles.

01:06:52.700 --> 01:06:55.020 align:middle line:90%
There should be no gap there.

01:06:55.020 --> 01:06:57.980 align:middle line:84%
Let's see, you
get two triangles.

01:06:57.980 --> 01:07:01.780 align:middle line:84%
And if you do g composed
with g composed with g,

01:07:01.780 --> 01:07:04.060 align:middle line:90%
you get four triangles.

01:07:04.060 --> 01:07:08.290 align:middle line:84%
So you need to think a little
bit to see why that is.

01:07:08.290 --> 01:07:10.510 align:middle line:84%
But actually, I think I
didn't need to show you.

01:07:10.510 --> 01:07:14.950 align:middle line:84%
But I was going to demonstrate
it on the graphing website.

01:07:14.950 --> 01:07:17.170 align:middle line:84%
But I think hopefully
the claim is clear,

01:07:17.170 --> 01:07:21.190 align:middle line:84%
that you can just with ReLUs,
make a triangle function.

01:07:21.190 --> 01:07:23.690 align:middle line:84%
And if you think about what
happens to the triangle function

01:07:23.690 --> 01:07:26.610 align:middle line:84%
when you compose it
with itself, you just

01:07:26.610 --> 01:07:29.810 align:middle line:84%
recursively get double as
many triangles each time.

01:07:29.810 --> 01:07:32.730 align:middle line:84%
So the point is that this
is an explicit construction

01:07:32.730 --> 01:07:36.370 align:middle line:84%
of a certain ReLU network, where
when you compose it with itself

01:07:36.370 --> 01:07:38.770 align:middle line:84%
many, many times,
you actually keep

01:07:38.770 --> 01:07:41.890 align:middle line:84%
doubling the number of
kinks in the function.

01:07:41.890 --> 01:07:46.790 align:middle line:84%
So it can really happen, at
least with this one example.

01:07:46.790 --> 01:07:49.570 align:middle line:90%


01:07:49.570 --> 01:07:51.490 align:middle line:84%
And then this is
actually just running

01:07:51.490 --> 01:07:54.410 align:middle line:84%
some numbers where
I'm like, let's

01:07:54.410 --> 01:07:56.290 align:middle line:84%
suppose that we take
that triangle function

01:07:56.290 --> 01:07:59.450 align:middle line:84%
and compose it with
itself 500 times,

01:07:59.450 --> 01:08:04.350 align:middle line:84%
that would give us at
most 2 to the 500 kinks.

01:08:04.350 --> 01:08:06.660 align:middle line:84%
No, it would actually
give us that many kinks.

01:08:06.660 --> 01:08:08.700 align:middle line:84%
It would really would
give us that many.

01:08:08.700 --> 01:08:12.360 align:middle line:90%


01:08:12.360 --> 01:08:15.260 align:middle line:84%
Actually, each g turns
out to have two layers.

01:08:15.260 --> 01:08:18.080 align:middle line:84%
So this would be
a 1,000-layer MLP.

01:08:18.080 --> 01:08:22.439 align:middle line:84%
And then if you compute using
the bound that we constructed,

01:08:22.439 --> 01:08:25.120 align:middle line:84%
if you took a three-layer
MLP, how wide would it

01:08:25.120 --> 01:08:27.899 align:middle line:84%
need to be in order to
fit the same function,

01:08:27.899 --> 01:08:30.520 align:middle line:84%
you'd find that it would
need to have almost 10

01:08:30.520 --> 01:08:34.520 align:middle line:84%
to the power 50 width in order
to fit that same function.

01:08:34.520 --> 01:08:36.540 align:middle line:90%
So this is the point.

01:08:36.540 --> 01:08:38.880 align:middle line:90%
This is the depth separation.

01:08:38.880 --> 01:08:40.640 align:middle line:84%
There's a function,
which is not even that

01:08:40.640 --> 01:08:41.819 align:middle line:90%
difficult to write down.

01:08:41.819 --> 01:08:48.779 align:middle line:84%
It's just a triangle or
self-composed many, many times.

01:08:48.779 --> 01:08:51.692 align:middle line:84%
And if you wanted to get a
shallow network to exactly

01:08:51.692 --> 01:08:53.359 align:middle line:84%
approximate that
function, it would need

01:08:53.359 --> 01:08:56.600 align:middle line:90%
to be very, very, very wide.

01:08:56.600 --> 01:08:58.660 align:middle line:84%
This is the reward for
coming to the lecture,

01:08:58.660 --> 01:09:02.080 align:middle line:84%
but this is also on
the first problem set.

01:09:02.080 --> 01:09:04.910 align:middle line:84%
So just maybe when you're
solving the problem set,

01:09:04.910 --> 01:09:08.790 align:middle line:90%
have a think about this.

01:09:08.790 --> 01:09:13.670 align:middle line:84%
OK, I want to now think
a bit more broadly.

01:09:13.670 --> 01:09:16.590 align:middle line:90%
What does this not say?

01:09:16.590 --> 01:09:18.910 align:middle line:84%
This depth separation
is suggesting

01:09:18.910 --> 01:09:20.430 align:middle line:84%
that there could
be a big advantage

01:09:20.430 --> 01:09:22.590 align:middle line:90%
to making things really deep.

01:09:22.590 --> 01:09:24.910 align:middle line:84%
But it's not telling us
that deep networks are

01:09:24.910 --> 01:09:27.630 align:middle line:84%
very easy to train,
and they might not be.

01:09:27.630 --> 01:09:30.950 align:middle line:84%
So it's not telling
us about optimization.

01:09:30.950 --> 01:09:34.612 align:middle line:84%
And it's not telling us
about generalization.

01:09:34.612 --> 01:09:37.029 align:middle line:84%
It's not saying that the very
deep networks would actually

01:09:37.029 --> 01:09:39.510 align:middle line:84%
perform well on some
particular data set

01:09:39.510 --> 01:09:42.350 align:middle line:90%
that we want to apply them to.

01:09:42.350 --> 01:09:45.830 align:middle line:84%
These are really just statements
thinking about approximation.

01:09:45.830 --> 01:09:47.770 align:middle line:84%
But this is not the
end of the story.

01:09:47.770 --> 01:09:50.050 align:middle line:84%
It's just one piece
in this puzzle.

01:09:50.050 --> 01:09:53.510 align:middle line:90%


01:09:53.510 --> 01:09:56.390 align:middle line:84%
Again, there's more
theorems of this nature.

01:09:56.390 --> 01:09:59.850 align:middle line:84%
And feel free to
look more into them.

01:09:59.850 --> 01:10:02.180 align:middle line:84%
The one at the bottom
is interesting.

01:10:02.180 --> 01:10:06.800 align:middle line:84%
Because a question you
could have is like, OK,

01:10:06.800 --> 01:10:14.140 align:middle line:84%
if increasing depth is so great,
maybe I'll just pick width 3

01:10:14.140 --> 01:10:19.700 align:middle line:84%
and just go really deep, and
this will work really well.

01:10:19.700 --> 01:10:24.100 align:middle line:84%
But what that paper shows is
that if the input space has

01:10:24.100 --> 01:10:27.540 align:middle line:84%
dimension n, you need
a width of at least n

01:10:27.540 --> 01:10:30.200 align:middle line:84%
to approximate any
function basically.

01:10:30.200 --> 01:10:33.540 align:middle line:84%
So there can be a minimum
width that you really need.

01:10:33.540 --> 01:10:35.620 align:middle line:84%
Anyway, there's some
further reading,

01:10:35.620 --> 01:10:39.900 align:middle line:84%
in case you're interested
to go deeper on this topic.

01:10:39.900 --> 01:10:47.060 align:middle line:84%
So in the last 15
minutes, I guess,

01:10:47.060 --> 01:10:50.620 align:middle line:84%
I want to think about some
more practical considerations.

01:10:50.620 --> 01:10:52.540 align:middle line:84%
Those were the
theoretical results.

01:10:52.540 --> 01:10:54.863 align:middle line:84%
Maybe they're
really interesting.

01:10:54.863 --> 01:10:57.280 align:middle line:84%
But if you're actually building
a machine learning system,

01:10:57.280 --> 01:11:01.570 align:middle line:84%
I don't know if you need to
think about those things.

01:11:01.570 --> 01:11:04.583 align:middle line:84%
Maybe you do, which
would be great.

01:11:04.583 --> 01:11:06.250 align:middle line:84%
But we're more just
trying to expose you

01:11:06.250 --> 01:11:08.433 align:middle line:84%
to these different
styles of thinking.

01:11:08.433 --> 01:11:10.850 align:middle line:84%
But now let's think a little
bit about some more practical

01:11:10.850 --> 01:11:12.110 align:middle line:90%
considerations.

01:11:12.110 --> 01:11:14.770 align:middle line:84%
So the trouble is in
practice, if you're actually

01:11:14.770 --> 01:11:17.065 align:middle line:84%
building a machine
learning system,

01:11:17.065 --> 01:11:18.690 align:middle line:84%
we were saying that
there's these three

01:11:18.690 --> 01:11:20.110 align:middle line:90%
pieces of the puzzle.

01:11:20.110 --> 01:11:24.530 align:middle line:84%
And isn't it nice if we can
neatly tackle them one by one?

01:11:24.530 --> 01:11:28.610 align:middle line:84%
But out there in the real
world, it's not like that.

01:11:28.610 --> 01:11:32.970 align:middle line:84%
And basically, if you're
training a system,

01:11:32.970 --> 01:11:34.770 align:middle line:84%
you have all of these
different issues

01:11:34.770 --> 01:11:36.450 align:middle line:90%
conflated with each other.

01:11:36.450 --> 01:11:38.950 align:middle line:84%
And if you've got a problem
and you're training,

01:11:38.950 --> 01:11:42.050 align:middle line:84%
you don't know if it's because
your neural network actually

01:11:42.050 --> 01:11:44.670 align:middle line:84%
can't approximate the thing that
you're trying to approximate.

01:11:44.670 --> 01:11:48.290 align:middle line:84%
You don't know if it's because
the training is failing.

01:11:48.290 --> 01:11:50.090 align:middle line:84%
Interestingly, you
would know if it

01:11:50.090 --> 01:11:51.667 align:middle line:84%
was because it's
not generalizing,

01:11:51.667 --> 01:11:54.250 align:middle line:84%
because you could just compute
the training error and the test

01:11:54.250 --> 01:11:56.083 align:middle line:84%
error, and you can see
if they're different.

01:11:56.083 --> 01:11:59.120 align:middle line:84%
So that one's
easier to diagnose.

01:11:59.120 --> 01:12:02.005 align:middle line:84%
But the point is that all of
these things are conflated.

01:12:02.005 --> 01:12:03.880 align:middle line:84%
And the other thing that
we were going to say

01:12:03.880 --> 01:12:08.000 align:middle line:84%
is, suppose that you work at
a large language model startup

01:12:08.000 --> 01:12:10.680 align:middle line:84%
and that your job is to
make the LLM training as

01:12:10.680 --> 01:12:13.280 align:middle line:84%
efficient as possible,
then you actually

01:12:13.280 --> 01:12:15.160 align:middle line:84%
really care about
this question--

01:12:15.160 --> 01:12:19.000 align:middle line:84%
should I make the network wider
or should I make it deeper?

01:12:19.000 --> 01:12:20.520 align:middle line:84%
For all those reasons
that we talked

01:12:20.520 --> 01:12:24.560 align:middle line:84%
about-- the computational
cost of training,

01:12:24.560 --> 01:12:27.160 align:middle line:84%
the computational cost
of inference, in terms

01:12:27.160 --> 01:12:30.480 align:middle line:90%
of making the model work well.

01:12:30.480 --> 01:12:32.480 align:middle line:84%
In a sense, it's the
most basic question

01:12:32.480 --> 01:12:33.560 align:middle line:84%
that you're thinking
about when you first

01:12:33.560 --> 01:12:34.727 align:middle line:90%
learn about neural networks.

01:12:34.727 --> 01:12:37.780 align:middle line:84%
But it's, to a large
extent, an unsolved problem.

01:12:37.780 --> 01:12:41.200 align:middle line:90%


01:12:41.200 --> 01:12:44.600 align:middle line:84%
And I wanted to talk a
little bit about this paper

01:12:44.600 --> 01:12:53.560 align:middle line:84%
by Kaplan and McCandlish from
2020, which, this figure, people

01:12:53.560 --> 01:12:55.780 align:middle line:84%
got T-shirts where they
just had this figure,

01:12:55.780 --> 01:12:58.730 align:middle line:84%
and then they would try
to spread the word about.

01:12:58.730 --> 01:13:01.270 align:middle line:84%
And this is the figure,
which is basically

01:13:01.270 --> 01:13:05.650 align:middle line:84%
saying as you make GPT
bigger, it gets better.

01:13:05.650 --> 01:13:10.910 align:middle line:84%
And the point is that,
in particular, the x-axis

01:13:10.910 --> 01:13:13.698 align:middle line:84%
is compute, the amount
of flops of compute,

01:13:13.698 --> 01:13:15.990 align:middle line:84%
the size of your data set,
and the number of parameters

01:13:15.990 --> 01:13:17.070 align:middle line:90%
in the model.

01:13:17.070 --> 01:13:19.830 align:middle line:84%
And the point is that the
test loss, the y-axis,

01:13:19.830 --> 01:13:24.070 align:middle line:84%
is always going down as you
scale any of these axes.

01:13:24.070 --> 01:13:28.230 align:middle line:84%
And the loss cannot go down
forever because you can't have

01:13:28.230 --> 01:13:29.350 align:middle line:90%
negative loss.

01:13:29.350 --> 01:13:30.690 align:middle line:90%
The loss has to go to 0.

01:13:30.690 --> 01:13:32.370 align:middle line:84%
So it has to
saturate eventually.

01:13:32.370 --> 01:13:34.970 align:middle line:84%
But back in 2020,
they were saying, hey,

01:13:34.970 --> 01:13:37.817 align:middle line:90%
the models keep getting better.

01:13:37.817 --> 01:13:39.150 align:middle line:90%
We should pay attention to this.

01:13:39.150 --> 01:13:43.990 align:middle line:84%
Because then all of these
advances in LLMs came.

01:13:43.990 --> 01:13:47.430 align:middle line:90%
ChatGPT was created and so on.

01:13:47.430 --> 01:13:51.550 align:middle line:84%
But interestingly, they plot
performance against parameters.

01:13:51.550 --> 01:13:53.910 align:middle line:84%
In this top plot,
it's not against width

01:13:53.910 --> 01:13:55.590 align:middle line:90%
and it's not against depth.

01:13:55.590 --> 01:13:59.100 align:middle line:84%
And the claim that
they have in that paper

01:13:59.100 --> 01:14:03.220 align:middle line:84%
is that within several
orders of magnitude,

01:14:03.220 --> 01:14:06.820 align:middle line:84%
the allocation of your compute
budget between width and depth

01:14:06.820 --> 01:14:07.940 align:middle line:90%
doesn't matter.

01:14:07.940 --> 01:14:10.500 align:middle line:84%
All that matters is the number
of parameters and number

01:14:10.500 --> 01:14:11.260 align:middle line:90%
of flops.

01:14:11.260 --> 01:14:15.100 align:middle line:84%
You can have a depth 20
network or a depth 10 network

01:14:15.100 --> 01:14:19.400 align:middle line:84%
with a commensurate number
of width, it doesn't matter.

01:14:19.400 --> 01:14:22.980 align:middle line:84%
So that's kind of interesting,
that you just run the experiment

01:14:22.980 --> 01:14:26.140 align:middle line:84%
and you see that width versus
depth, within a large range,

01:14:26.140 --> 01:14:27.320 align:middle line:90%
seems not to matter.

01:14:27.320 --> 01:14:31.900 align:middle line:84%
And they exemplify that
with this lower plot.

01:14:31.900 --> 01:14:34.620 align:middle line:84%
And the point is that once
they subtract the embedding

01:14:34.620 --> 01:14:37.680 align:middle line:84%
parameters-- so if we just focus
in on the bottom-right plot,

01:14:37.680 --> 01:14:42.420 align:middle line:84%
which is the model with the
embedding layers not counted--

01:14:42.420 --> 01:14:45.080 align:middle line:84%
roughly you get very
similar performance.

01:14:45.080 --> 01:14:48.400 align:middle line:84%
I think what basically they're
saying-- beyond depth 6,

01:14:48.400 --> 01:14:51.080 align:middle line:84%
it doesn't matter what the
width is or what the depth is.

01:14:51.080 --> 01:14:53.380 align:middle line:84%
The curves converge
to each other.

01:14:53.380 --> 01:14:56.770 align:middle line:84%
And all that matters is scaling
up number of parameters.

01:14:56.770 --> 01:15:01.550 align:middle line:84%
So this is just a very,
very practical perspective,

01:15:01.550 --> 01:15:08.210 align:middle line:84%
but from some very
careful experimentalists

01:15:08.210 --> 01:15:10.610 align:middle line:84%
about how things actually
work in transformers.

01:15:10.610 --> 01:15:13.370 align:middle line:84%
And you see it's
basically saying--

01:15:13.370 --> 01:15:16.390 align:middle line:84%
within a large range, width
versus depth doesn't matter;

01:15:16.390 --> 01:15:18.010 align:middle line:84%
all you care about
is parameters.

01:15:18.010 --> 01:15:23.170 align:middle line:84%
With that said,
there's something

01:15:23.170 --> 01:15:25.250 align:middle line:84%
I wanted to draw
attention to, which

01:15:25.250 --> 01:15:28.010 align:middle line:90%
is this idea of confounders.

01:15:28.010 --> 01:15:29.930 align:middle line:84%
And there's another
paper which people

01:15:29.930 --> 01:15:32.050 align:middle line:84%
are also excited about,
which people refer

01:15:32.050 --> 01:15:34.710 align:middle line:90%
to as chinchilla scaling rules.

01:15:34.710 --> 01:15:37.090 align:middle line:84%
The names keep
getting stranger--

01:15:37.090 --> 01:15:37.958 align:middle line:90%
[LAUGHTER]

01:15:37.958 --> 01:15:38.750 align:middle line:90%
--for these things.

01:15:38.750 --> 01:15:43.970 align:middle line:84%
But basically, they're
interested in the same sort

01:15:43.970 --> 01:15:47.250 align:middle line:84%
of questions as the previous
paper I just showed you.

01:15:47.250 --> 01:15:49.730 align:middle line:84%
But what they do
is they question

01:15:49.730 --> 01:15:52.200 align:middle line:84%
some of the results
in that earlier paper.

01:15:52.200 --> 01:15:54.200 align:middle line:84%
And one of the things
that they suggest

01:15:54.200 --> 01:15:56.680 align:middle line:84%
is that if you use a
different learning rate

01:15:56.680 --> 01:15:59.160 align:middle line:84%
schedule for the
training, you can

01:15:59.160 --> 01:16:04.120 align:middle line:84%
get slightly different,
qualitatively

01:16:04.120 --> 01:16:05.980 align:middle line:90%
different conclusions.

01:16:05.980 --> 01:16:08.160 align:middle line:84%
So if you just had one
learning rate schedule

01:16:08.160 --> 01:16:11.060 align:middle line:84%
and now you do a different
learning rate schedule,

01:16:11.060 --> 01:16:14.640 align:middle line:90%
you draw different conclusions.

01:16:14.640 --> 01:16:17.440 align:middle line:84%
And this is just to say
that precisely answering

01:16:17.440 --> 01:16:21.800 align:middle line:84%
any of these questions
is very difficult,

01:16:21.800 --> 01:16:24.880 align:middle line:84%
because there are so many
parts of the training pipeline

01:16:24.880 --> 01:16:26.400 align:middle line:90%
that we don't understand.

01:16:26.400 --> 01:16:28.280 align:middle line:84%
So you may have one
experimental setup

01:16:28.280 --> 01:16:31.720 align:middle line:84%
and you say, hey, width
versus depth doesn't matter.

01:16:31.720 --> 01:16:34.080 align:middle line:84%
And then you change some
detail about how you actually

01:16:34.080 --> 01:16:34.940 align:middle line:90%
train the network.

01:16:34.940 --> 01:16:37.520 align:middle line:84%
And now it's like
oh, no, scaling width

01:16:37.520 --> 01:16:40.420 align:middle line:84%
is a lot better because
I fixed that issue.

01:16:40.420 --> 01:16:43.360 align:middle line:84%
And I would call that a
confounding issue, which is just

01:16:43.360 --> 01:16:45.440 align:middle line:84%
some aspect of the training
that we don't fully

01:16:45.440 --> 01:16:49.240 align:middle line:84%
understand, but is having
a bearing on the results.

01:16:49.240 --> 01:16:52.290 align:middle line:84%
So this is just what
I want to point out,

01:16:52.290 --> 01:16:55.190 align:middle line:84%
is that really resolving
these questions

01:16:55.190 --> 01:16:58.910 align:middle line:84%
experimentally is really
hard, because there can be

01:16:58.910 --> 01:17:00.610 align:middle line:90%
so many confounding variables.

01:17:00.610 --> 01:17:03.450 align:middle line:84%
And resolving them theoretically
is also really hard,

01:17:03.450 --> 01:17:06.750 align:middle line:84%
because you have to work out
how to think about these things.

01:17:06.750 --> 01:17:09.497 align:middle line:84%
And a lot of these problems are
unsolved because of the fact

01:17:09.497 --> 01:17:11.330 align:middle line:84%
that it's kind of hard
to figure things out.

01:17:11.330 --> 01:17:14.190 align:middle line:90%


01:17:14.190 --> 01:17:19.110 align:middle line:90%
So in conclusion, the summary.

01:17:19.110 --> 01:17:21.930 align:middle line:84%
Oh, yeah, there's still
another slide after this one.

01:17:21.930 --> 01:17:25.870 align:middle line:84%
When you leave the lecture
hall, try to do so quietly.

01:17:25.870 --> 01:17:28.310 align:middle line:90%
Someone suggested that.

01:17:28.310 --> 01:17:31.170 align:middle line:84%
OK, so the summary
of the lecture is--

01:17:31.170 --> 01:17:34.070 align:middle line:90%


01:17:34.070 --> 01:17:37.790 align:middle line:84%
a very wide shallow neural
net, so a three-layer

01:17:37.790 --> 01:17:41.990 align:middle line:84%
MLP that's wide enough
can fit any function

01:17:41.990 --> 01:17:44.510 align:middle line:84%
within some class
of functions, quite

01:17:44.510 --> 01:17:46.010 align:middle line:90%
a broad class of functions.

01:17:46.010 --> 01:17:49.360 align:middle line:90%


01:17:49.360 --> 01:17:51.040 align:middle line:84%
Then, in the second
half of the lecture,

01:17:51.040 --> 01:17:54.300 align:middle line:84%
we said that deeper networks can
fit certain kinds of functions

01:17:54.300 --> 01:17:57.820 align:middle line:84%
with many fewer neurons because
of their compositionality.

01:17:57.820 --> 01:18:00.360 align:middle line:84%
And we called that a kind
of depth separation result.

01:18:00.360 --> 01:18:02.220 align:middle line:84%
So that was something
interesting.

01:18:02.220 --> 01:18:12.820 align:middle line:84%
And then the overall message
that I am hoping to impart

01:18:12.820 --> 01:18:17.940 align:middle line:84%
is basically, it's unclear how
important these results are.

01:18:17.940 --> 01:18:20.620 align:middle line:84%
It's unclear how they
interact with optimization

01:18:20.620 --> 01:18:24.020 align:middle line:84%
and with generalization, but
it's still really interesting

01:18:24.020 --> 01:18:27.930 align:middle line:84%
to think about them and to have
that as part of your toolkit.

01:18:27.930 --> 01:18:30.180 align:middle line:84%
If you're thinking about
your machine learning system,

01:18:30.180 --> 01:18:32.460 align:middle line:84%
a basic question
is, can the network

01:18:32.460 --> 01:18:34.020 align:middle line:84%
architecture that
I have actually

01:18:34.020 --> 01:18:37.460 align:middle line:84%
approximate the function that
I'm asking it to approximate?

01:18:37.460 --> 01:18:41.260 align:middle line:90%
That's quite a basic question.

01:18:41.260 --> 01:18:43.540 align:middle line:84%
And I want to just
conclude or to wrap up

01:18:43.540 --> 01:18:46.220 align:middle line:84%
with a little preview for
what's to come in the course.

01:18:46.220 --> 01:18:48.530 align:middle line:90%
So I have a question.

01:18:48.530 --> 01:18:51.690 align:middle line:84%
And it's, suppose that we
have two machine learning

01:18:51.690 --> 01:18:55.290 align:middle line:84%
problems-- one is
an audio problem

01:18:55.290 --> 01:18:58.630 align:middle line:84%
and one is an image
classification problem.

01:18:58.630 --> 01:19:01.850 align:middle line:84%
So the first little
picture is like a waveform.

01:19:01.850 --> 01:19:05.090 align:middle line:84%
And then we're transcribing
it to say the word "hello."

01:19:05.090 --> 01:19:08.050 align:middle line:84%
And then the second picture is
like a photograph of a person.

01:19:08.050 --> 01:19:12.810 align:middle line:84%
And then we're classifying
that as, that's a human.

01:19:12.810 --> 01:19:15.550 align:middle line:84%
And then the thought
experiment is, in both cases,

01:19:15.550 --> 01:19:21.030 align:middle line:84%
we could flatten those little
pieces of data into vectors.

01:19:21.030 --> 01:19:23.250 align:middle line:84%
OK, maybe the
waveform is a vector.

01:19:23.250 --> 01:19:25.570 align:middle line:84%
And then we can flatten
that image of a person

01:19:25.570 --> 01:19:27.170 align:middle line:90%
into another vector.

01:19:27.170 --> 01:19:31.650 align:middle line:84%
And then we can just
apply MLP to both.

01:19:31.650 --> 01:19:33.650 align:middle line:84%
We can take a
multilayer perceptron

01:19:33.650 --> 01:19:37.510 align:middle line:84%
with really nonlinearity and
try to fit the first problem.

01:19:37.510 --> 01:19:39.770 align:middle line:84%
And then we can take
another MLP and try

01:19:39.770 --> 01:19:41.690 align:middle line:90%
to fit the second problem.

01:19:41.690 --> 01:19:45.640 align:middle line:84%
And the argument is, OK, the
MLP is a universal function

01:19:45.640 --> 01:19:49.520 align:middle line:84%
approximator, so
this is a great idea.

01:19:49.520 --> 01:19:53.786 align:middle line:84%
And so I want to ask, do people
think that is a good idea?

01:19:53.786 --> 01:19:54.640 align:middle line:90%
AUDIENCE: No!

01:19:54.640 --> 01:19:55.515 align:middle line:90%
JEREMY BERNSTEIN: No?

01:19:55.515 --> 01:19:58.040 align:middle line:90%
[LAUGHS] Why not?

01:19:58.040 --> 01:20:00.760 align:middle line:90%
But why not?

01:20:00.760 --> 01:20:05.040 align:middle line:84%
AUDIENCE: You're not guaranteed
to find the approximator.

01:20:05.040 --> 01:20:07.265 align:middle line:84%
It exists, but you
may not find it.

01:20:07.265 --> 01:20:09.140 align:middle line:84%
JEREMY BERNSTEIN: Yeah,
that's one objection.

01:20:09.140 --> 01:20:11.500 align:middle line:84%
That's a great point,
that it could exist,

01:20:11.500 --> 01:20:12.820 align:middle line:90%
but we might not find it.

01:20:12.820 --> 01:20:16.240 align:middle line:90%


01:20:16.240 --> 01:20:19.640 align:middle line:84%
AUDIENCE: I'd say efficiency,
because although it

01:20:19.640 --> 01:20:21.720 align:middle line:84%
is a universal
function approximation

01:20:21.720 --> 01:20:25.320 align:middle line:84%
in the [INAUDIBLE], it says
you may or may not find

01:20:25.320 --> 01:20:27.760 align:middle line:90%
that function in that space.

01:20:27.760 --> 01:20:30.260 align:middle line:84%
To find it, you may need
to tune certain parameters.

01:20:30.260 --> 01:20:32.660 align:middle line:84%
It may take a while to find
what you're looking for it.

01:20:32.660 --> 01:20:37.423 align:middle line:84%
So after a certain cutoff,
it may not be feasible.

01:20:37.423 --> 01:20:38.840 align:middle line:84%
JEREMY BERNSTEIN:
The objection is

01:20:38.840 --> 01:20:41.100 align:middle line:84%
it could be inefficient
to use an MLP.

01:20:41.100 --> 01:20:44.310 align:middle line:84%
There might be potentially
another the architecture

01:20:44.310 --> 01:20:47.310 align:middle line:90%
that could be more efficient.

01:20:47.310 --> 01:20:50.090 align:middle line:84%
And in particular, it could be
adapted to the type of data.

01:20:50.090 --> 01:20:51.053 align:middle line:90%
Sorry?

01:20:51.053 --> 01:20:52.470 align:middle line:84%
AUDIENCE: I was
just going to say,

01:20:52.470 --> 01:20:56.990 align:middle line:84%
an MLP doesn't take advantage of
the structure in the input data.

01:20:56.990 --> 01:20:59.430 align:middle line:84%
So with an image, there's
an inherent two-dimensional

01:20:59.430 --> 01:21:00.370 align:middle line:90%
structure.

01:21:00.370 --> 01:21:03.862 align:middle line:84%
And with a waveform, there's an
inherent sequential structure.

01:21:03.862 --> 01:21:05.570 align:middle line:84%
JEREMY BERNSTEIN:
That's a great comment.

01:21:05.570 --> 01:21:09.390 align:middle line:84%
So the comment is that
the two different problems

01:21:09.390 --> 01:21:11.090 align:middle line:84%
have different
structure in the data.

01:21:11.090 --> 01:21:13.598 align:middle line:84%
And perhaps we want to
adapt the model family.

01:21:13.598 --> 01:21:15.890 align:middle line:84%
And that's what I thought
when I was writing the slide.

01:21:15.890 --> 01:21:17.530 align:middle line:84%
And then I was like,
but wait a minute,

01:21:17.530 --> 01:21:20.170 align:middle line:84%
now we're just applying
transformers to everything.

01:21:20.170 --> 01:21:22.130 align:middle line:84%
So is that counter
to that point?

01:21:22.130 --> 01:21:24.810 align:middle line:84%
And then it was like,
no, actually it's not.

01:21:24.810 --> 01:21:27.430 align:middle line:84%
Because even if you apply a
transformer to audio or you

01:21:27.430 --> 01:21:30.310 align:middle line:84%
apply it to images, actually
the preprocessing layer

01:21:30.310 --> 01:21:32.830 align:middle line:84%
at the beginning of the
network is very different.

01:21:32.830 --> 01:21:36.350 align:middle line:84%
So if it's images, you have
a special 2D patch kind

01:21:36.350 --> 01:21:37.570 align:middle line:90%
of representation.

01:21:37.570 --> 01:21:39.790 align:middle line:84%
And if it's audio, I
assume they have some kind

01:21:39.790 --> 01:21:41.790 align:middle line:90%
of better audio representation.

01:21:41.790 --> 01:21:44.740 align:middle line:84%
So even if you just throw
transformers at everything,

01:21:44.740 --> 01:21:47.060 align:middle line:84%
actually this
comment that you do

01:21:47.060 --> 01:21:50.460 align:middle line:84%
adapt the architecture a little
bit based on what the data is.

01:21:50.460 --> 01:21:55.260 align:middle line:84%
So yeah, I really agree
with that comment.

01:21:55.260 --> 01:21:56.820 align:middle line:90%
But it was a big question.

01:21:56.820 --> 01:21:59.240 align:middle line:84%
People even recently have been
thinking about that-- wait,

01:21:59.240 --> 01:22:00.520 align:middle line:90%
are transformers all we need?

01:22:00.520 --> 01:22:02.760 align:middle line:84%
Maybe we don't need
to model the data.

01:22:02.760 --> 01:22:04.500 align:middle line:84%
But I think really,
you still need

01:22:04.500 --> 01:22:07.540 align:middle line:90%
to model the data little bit.

01:22:07.540 --> 01:22:11.660 align:middle line:84%
And also transformers, you can
ask how data hungry are they.

01:22:11.660 --> 01:22:13.320 align:middle line:84%
Maybe there's a
better model family

01:22:13.320 --> 01:22:15.750 align:middle line:84%
which needs much less
data in order to learn.

01:22:15.750 --> 01:22:17.500 align:middle line:84%
So that was the final
thought to leave you

01:22:17.500 --> 01:22:19.260 align:middle line:84%
with is, perhaps
we want to match

01:22:19.260 --> 01:22:22.560 align:middle line:84%
the architecture to the problem
that we're trying to solve,

01:22:22.560 --> 01:22:24.080 align:middle line:90%
or the structure of the data.

01:22:24.080 --> 01:22:26.460 align:middle line:84%
Perhaps it could be more
computationally efficient.

01:22:26.460 --> 01:22:29.860 align:middle line:84%
Perhaps it could make it
easier to find the approximator

01:22:29.860 --> 01:22:32.100 align:middle line:90%
that we're looking for.

01:22:32.100 --> 01:22:33.560 align:middle line:84%
OK, that's the end
of the lecture.

01:22:33.560 --> 01:22:35.630 align:middle line:90%
Thank you, everyone.

01:22:35.630 --> 01:22:42.000 align:middle line:90%